mirror of
https://github.com/agentscope-ai/ReMe.git
synced 2026-10-10 03:30:56 +00:00
refactor(auto_fin): replace similarity with direction classification for historical events (#395)
- Add AutoFinHistoricalDirectionReference model to classify historical events by direction - Remove AutoFinHistoricalSimilarity and related similarity score usage - Update AutoFinMarketSelection to handle same and opposite direction event lists - Adjust AutoFinMarketStep to calculate forecasts based on equal weights and direction signs - Change market.yaml instructions to require direction classification instead of similarity scoring - Modify tests to reflect direction-based classification and verify uniqueness across direction groups - Improve DingTalkWaitStep to support reconnect on server request with proper disconnect reason handling
This commit is contained in:
parent
0522135791
commit
2f79977df0
7 changed files with 200 additions and 71 deletions
|
|
@ -13,8 +13,8 @@ from .auto_fin import (
|
|||
AutoFinFutureReturnPoint,
|
||||
AutoFinHistoricalEvent,
|
||||
AutoFinHistoricalEventReference,
|
||||
AutoFinHistoricalDirectionReference,
|
||||
AutoFinHistoricalMatch,
|
||||
AutoFinHistoricalSimilarity,
|
||||
AutoFinMarketSelection,
|
||||
AutoFinMarketSample,
|
||||
AutoFinReportOutput,
|
||||
|
|
@ -54,8 +54,8 @@ __all__ = [
|
|||
"AutoFinFutureReturnPoint",
|
||||
"AutoFinHistoricalEvent",
|
||||
"AutoFinHistoricalEventReference",
|
||||
"AutoFinHistoricalDirectionReference",
|
||||
"AutoFinHistoricalMatch",
|
||||
"AutoFinHistoricalSimilarity",
|
||||
"AutoFinMarketSelection",
|
||||
"AutoFinMarketSample",
|
||||
"AutoFinReportOutput",
|
||||
|
|
|
|||
|
|
@ -235,44 +235,44 @@ class AutoFinEtfHistoricalResearch(AutoFinModel):
|
|||
return self
|
||||
|
||||
|
||||
class AutoFinHistoricalSimilarity(AutoFinModel):
|
||||
"""One similarity judgment returned by the Market Agent."""
|
||||
class AutoFinHistoricalDirectionReference(AutoFinModel):
|
||||
"""One direction-classified historical event returned by the Market Agent."""
|
||||
|
||||
reason: str
|
||||
news_id: str
|
||||
similarity: float
|
||||
|
||||
@model_validator(mode="after")
|
||||
def non_empty_values(self) -> "AutoFinHistoricalSimilarity":
|
||||
"""Reject a similarity judgment without source identity or rationale."""
|
||||
def non_empty_values(self) -> "AutoFinHistoricalDirectionReference":
|
||||
"""Reject a direction judgment without source identity or rationale."""
|
||||
self.reason = self.reason.strip()
|
||||
self.news_id = self.news_id.strip()
|
||||
if not self.reason or not self.news_id:
|
||||
raise ValueError("historical similarity reason and news ID must be non-empty")
|
||||
raise ValueError("historical direction reason and news ID must be non-empty")
|
||||
return self
|
||||
|
||||
|
||||
class AutoFinMarketSelection(AutoFinModel):
|
||||
"""Historical similarities returned by the Market Agent."""
|
||||
"""Same- and opposite-direction historical events returned by the Market Agent."""
|
||||
|
||||
matched_historical_events: list[AutoFinHistoricalSimilarity] = Field(default_factory=list)
|
||||
same_direction_events: list[AutoFinHistoricalDirectionReference] = Field(default_factory=list)
|
||||
opposite_direction_events: list[AutoFinHistoricalDirectionReference] = Field(default_factory=list)
|
||||
|
||||
@model_validator(mode="after")
|
||||
def unique_historical_news(self) -> "AutoFinMarketSelection":
|
||||
"""Reject duplicate similarity judgments."""
|
||||
news_ids = [event.news_id for event in self.matched_historical_events]
|
||||
"""Reject news IDs repeated within or across direction groups."""
|
||||
news_ids = [event.news_id for event in (*self.same_direction_events, *self.opposite_direction_events)]
|
||||
if len(news_ids) != len(set(news_ids)):
|
||||
raise ValueError("matched historical event news IDs must be unique")
|
||||
raise ValueError("direction-classified historical event news IDs must be unique")
|
||||
return self
|
||||
|
||||
|
||||
class AutoFinHistoricalMatch(AutoFinModel):
|
||||
"""One historical event selected for the weighted forecast."""
|
||||
"""One direction-classified historical event used by the equal-weight forecast."""
|
||||
|
||||
reason: str
|
||||
news_id: str
|
||||
event_time: ShanghaiDateTime
|
||||
similarity: float
|
||||
direction: Literal["same", "opposite"]
|
||||
weight: float = Field(ge=0.0, le=1.0)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -17,7 +17,7 @@ from ._base import AutoFinStep, _write
|
|||
|
||||
@R.register("auto_fin_market_step")
|
||||
class AutoFinMarketStep(AutoFinStep):
|
||||
"""Collect similarity judgments and calculate one ETF forecast."""
|
||||
"""Classify historical event directions and calculate one ETF forecast."""
|
||||
|
||||
@staticmethod
|
||||
def _calculate_analysis(
|
||||
|
|
@ -27,27 +27,24 @@ class AutoFinMarketStep(AutoFinStep):
|
|||
) -> AutoFinSelectedEtfAnalysis:
|
||||
"""Build all deterministic market fields from Agent-selected news IDs."""
|
||||
history_by_news_id = {event.news_id: event for event in history.historical_events}
|
||||
unknown_news_ids = {
|
||||
match.news_id for match in selection.matched_historical_events if match.news_id not in history_by_news_id
|
||||
}
|
||||
selected = [
|
||||
*((match, "same", 1.0) for match in selection.same_direction_events),
|
||||
*((match, "opposite", -1.0) for match in selection.opposite_direction_events),
|
||||
]
|
||||
unknown_news_ids = {match.news_id for match, _, _ in selected if match.news_id not in history_by_news_id}
|
||||
if unknown_news_ids:
|
||||
raise ValueError(f"Market Agent referenced unknown historical news: {sorted(unknown_news_ids)}")
|
||||
|
||||
selected = [
|
||||
(match, min(1.0, max(-1.0, match.similarity)))
|
||||
for match in selection.matched_historical_events
|
||||
if min(1.0, max(-1.0, match.similarity)) != 0
|
||||
]
|
||||
total_similarity = sum(abs(similarity) for _, similarity in selected)
|
||||
weight = 1.0 / len(selected) if selected else 0.0
|
||||
matches = [
|
||||
{
|
||||
"reason": match.reason,
|
||||
"news_id": match.news_id,
|
||||
"event_time": history_by_news_id[match.news_id].event_time,
|
||||
"similarity": similarity,
|
||||
"weight": abs(similarity) / total_similarity,
|
||||
"direction": direction,
|
||||
"weight": weight,
|
||||
}
|
||||
for match, similarity in selected
|
||||
for match, direction, _ in selected
|
||||
]
|
||||
|
||||
returns = []
|
||||
|
|
@ -55,23 +52,19 @@ class AutoFinMarketStep(AutoFinStep):
|
|||
has_direction_conflict = False
|
||||
for horizon in range(1, 11):
|
||||
available = []
|
||||
for match, similarity in selected:
|
||||
for match, _, direction_coefficient in selected:
|
||||
event = history_by_news_id[match.news_id]
|
||||
point = next((point for point in event.future_returns if point.horizon == horizon), None)
|
||||
if point is not None:
|
||||
direction = 1.0 if similarity > 0 else -1.0
|
||||
available.append((abs(similarity), direction * point.cumulative_return))
|
||||
available.append(direction_coefficient * point.cumulative_return)
|
||||
if not available:
|
||||
has_missing_horizon = True
|
||||
expected_return = None
|
||||
else:
|
||||
horizon_similarity = sum(similarity for similarity, _ in available)
|
||||
expected_return = (
|
||||
sum(similarity * cumulative_return for similarity, cumulative_return in available)
|
||||
/ horizon_similarity
|
||||
expected_return = sum(available) / len(available)
|
||||
has_direction_conflict |= any(value > 0 for value in available) and any(
|
||||
value < 0 for value in available
|
||||
)
|
||||
values = [cumulative_return for _, cumulative_return in available]
|
||||
has_direction_conflict |= any(value > 0 for value in values) and any(value < 0 for value in values)
|
||||
returns.append({"horizon": horizon, "expected_return": expected_return})
|
||||
|
||||
positive_returns = [point for point in returns if (point["expected_return"] or 0) > 0]
|
||||
|
|
@ -132,7 +125,7 @@ class AutoFinMarketStep(AutoFinStep):
|
|||
selection_path = (
|
||||
self.workspace_path / "resource" / str(self._required("auto_fin_date")) / f"{resource_name}_output.json"
|
||||
)
|
||||
self.logger.warning(f"[{self.name}] skip similarity Agent for {item.etf_code}: no valid history")
|
||||
self.logger.warning(f"[{self.name}] skip direction Agent for {item.etf_code}: no valid history")
|
||||
analysis = self._calculate_analysis(item, history, selection)
|
||||
_write(
|
||||
selection_path,
|
||||
|
|
|
|||
|
|
@ -1,5 +1,6 @@
|
|||
market_user: |
|
||||
你只负责判断历史事件与当前事件的相似度,不计算收益、不生成预测、不总结报告。
|
||||
你只负责筛选与当前事件机制可比的历史事件,并判断影响方向相同还是相反;不计算收益、
|
||||
不生成预测、不总结报告。
|
||||
ETF:{etf_code}({etf_name})
|
||||
分析截止时间:{decision_at}
|
||||
|
||||
|
|
@ -11,25 +12,33 @@ market_user: |
|
|||
|
||||
必须遵守:
|
||||
1. 读取历史文件,只从 historical_events 中选择与当前事件相似的事件。
|
||||
2. 综合事件类型、关键实体、传导机制和影响方向判断 similarity,范围为 [-1, 1]。正值表示
|
||||
机制和影响方向相似,负值表示机制可比但影响方向相反,0 表示没有有效关系。例如当前事件是
|
||||
黄金涨价、历史事件是黄金降价时,可以返回负值。程序会将越界值截断到 [-1, 1]。
|
||||
2. 综合事件类型、关键实体、传导机制和影响方向进行判断:
|
||||
- 机制可比且影响方向相同,放入 same_direction_events;
|
||||
- 机制可比但影响方向相反,放入 opposite_direction_events;
|
||||
- 没有有效关系,不要返回。
|
||||
这里的“相同/相反”是历史事件相对当前事件的影响方向,不是简单判断新闻利好或利空。例如
|
||||
当前事件是黄金涨价、历史事件是黄金降价时,应放入 opposite_direction_events。
|
||||
3. 只根据历史事件的 event_time、event_title、event_content 和 reason 判断相似性;不要依据
|
||||
market_entry 或 future_returns 选择事件,避免使用事后行情影响相似度判断。
|
||||
4. reason 简洁说明相似之处,news_id 必须从历史文件逐字复制,严禁编造。
|
||||
5. 不需要返回 ETF、事件时间、权重、预测、持有天数、代码、总结或 limitations;程序会校验
|
||||
news_id,并完成所有计算。没有相似事件时返回空列表。
|
||||
6. 最终只生成 matched_historical_events;每项字段顺序为 reason、news_id、similarity。
|
||||
JSON 示例:
|
||||
news_id,并完成所有计算。没有相似事件时两组都返回空列表。
|
||||
6. 同一 news_id 只能出现一次,不能同时放入两组。每项只包含 reason、news_id。
|
||||
7. 最终只生成 same_direction_events 和 opposite_direction_events。JSON 示例:
|
||||
```json
|
||||
{{
|
||||
"matched_historical_events": [
|
||||
"same_direction_events": [
|
||||
{{
|
||||
"reason": "事件类型、关键实体、传导机制和影响方向相似",
|
||||
"news_id": "20260601100000_a3f8",
|
||||
"similarity": 0.86
|
||||
"news_id": "20260601100000_a3f8"
|
||||
}}
|
||||
],
|
||||
"opposite_direction_events": [
|
||||
{{
|
||||
"reason": "传导机制相似,但影响方向相反",
|
||||
"news_id": "20260501100000_b4c9"
|
||||
}}
|
||||
]
|
||||
}}
|
||||
```
|
||||
7. 最终只生成上述 JSON 对象,不附加解释或 Markdown 正文。
|
||||
8. 最终只生成上述 JSON 对象,不附加解释或 Markdown 正文。
|
||||
|
|
|
|||
|
|
@ -71,7 +71,7 @@ class DingTalkWaitStep(BaseStep):
|
|||
)
|
||||
workers = [asyncio.create_task(self._worker(queue, locks, sessions, handler)) for _ in range(self.worker_count)]
|
||||
try:
|
||||
await self._run_client(client, self.context.stop_event)
|
||||
await self._run_with_reconnect(client, self.context.stop_event)
|
||||
finally:
|
||||
for worker in workers:
|
||||
worker.cancel()
|
||||
|
|
@ -169,9 +169,24 @@ class DingTalkWaitStep(BaseStep):
|
|||
)
|
||||
raise
|
||||
|
||||
async def _run_with_reconnect(self, client, stop_event: asyncio.Event) -> None:
|
||||
"""Reconnect cleanly when DingTalk rotates an otherwise healthy connection."""
|
||||
while not stop_event.is_set():
|
||||
disconnect_reason = await self._run_client(client, stop_event)
|
||||
if disconnect_reason is None:
|
||||
return
|
||||
self.logger.info(
|
||||
f"[{self.name}] DingTalk server requested reconnect reason={disconnect_reason!r}; "
|
||||
"reconnecting in 1.0s",
|
||||
)
|
||||
try:
|
||||
await asyncio.wait_for(stop_event.wait(), timeout=1.0)
|
||||
except asyncio.TimeoutError:
|
||||
pass
|
||||
|
||||
@staticmethod
|
||||
async def _run_client(client, stop_event: asyncio.Event) -> None:
|
||||
"""Run one cancellable WebSocket connection; BackgroundJob owns retries."""
|
||||
async def _run_client(client, stop_event: asyncio.Event) -> str | None:
|
||||
"""Run one connection and return a server-requested disconnect reason."""
|
||||
import websockets # pylint: disable=import-outside-toplevel
|
||||
|
||||
client.pre_start()
|
||||
|
|
@ -179,6 +194,7 @@ class DingTalkWaitStep(BaseStep):
|
|||
if not connection:
|
||||
raise ConnectionError("DingTalk open connection failed")
|
||||
uri = f'{connection["endpoint"]}?ticket={quote_plus(connection["ticket"])}'
|
||||
disconnect_reason = None
|
||||
async with websockets.connect(uri) as websocket:
|
||||
client.websocket = websocket
|
||||
keepalive = asyncio.create_task(client.keepalive(websocket))
|
||||
|
|
@ -190,11 +206,26 @@ class DingTalkWaitStep(BaseStep):
|
|||
stopper = asyncio.create_task(close_when_stopped())
|
||||
try:
|
||||
async for raw_message in websocket:
|
||||
if await client.route_message(json.loads(raw_message)) == client.TAG_DISCONNECT:
|
||||
message = json.loads(raw_message)
|
||||
if await client.route_message(message) == client.TAG_DISCONNECT:
|
||||
data = message.get("data", {})
|
||||
if isinstance(data, str):
|
||||
try:
|
||||
data = json.loads(data)
|
||||
except json.JSONDecodeError:
|
||||
data = {}
|
||||
reason = data.get("reason") if isinstance(data, dict) else None
|
||||
disconnect_reason = (
|
||||
reason.strip() if isinstance(reason, str) and reason.strip() else "unspecified"
|
||||
)
|
||||
await websocket.close()
|
||||
break
|
||||
finally:
|
||||
for task in (stopper, keepalive):
|
||||
task.cancel()
|
||||
await asyncio.gather(stopper, keepalive, return_exceptions=True)
|
||||
if not stop_event.is_set():
|
||||
raise ConnectionError("DingTalk WebSocket closed")
|
||||
if stop_event.is_set():
|
||||
return None
|
||||
if disconnect_reason is not None:
|
||||
return disconnect_reason
|
||||
raise ConnectionError("DingTalk WebSocket closed unexpectedly")
|
||||
|
|
|
|||
|
|
@ -375,9 +375,10 @@ class _Agent(BaseAgentWrapper):
|
|||
elif schema is AutoFinMarketSelection:
|
||||
assert "ETF:159018.SZ(油气ETF)" in task
|
||||
assert "[2026-07-23T16:00:00] 原油供应中断" in task
|
||||
assert "只负责判断历史事件与当前事件的相似度" in task
|
||||
assert "判断影响方向相同还是相反" in task
|
||||
assert "不要依据" in task
|
||||
assert "程序会校验" in task
|
||||
assert "每项只包含 reason、news_id" in task
|
||||
assert "$tushare-data" not in task
|
||||
history_path = Path(
|
||||
next(line.strip() for line in task.splitlines() if line.strip().endswith("_output.json")),
|
||||
|
|
@ -388,13 +389,13 @@ class _Agent(BaseAgentWrapper):
|
|||
assert history["historical_events"][0]["event_title"] == "历史供应中断"
|
||||
assert len(history["historical_events"][0]["future_returns"]) == 10
|
||||
value = {
|
||||
"matched_historical_events": [
|
||||
"same_direction_events": [
|
||||
{
|
||||
"reason": "供应中断的事件类型和传导机制相同",
|
||||
"news_id": history["historical_events"][0]["news_id"],
|
||||
"similarity": 1.2,
|
||||
},
|
||||
],
|
||||
"opposite_direction_events": [],
|
||||
}
|
||||
elif schema is AutoFinReportOutput:
|
||||
assert "不重新搜索新闻" in task
|
||||
|
|
@ -550,7 +551,7 @@ async def test_four_step_pipeline_writes_plain_markdown_and_cleans_temporary_dat
|
|||
analysis = detail["market_analysis"]
|
||||
assert analysis["matched_historical_events"][0]["weight"] == 1.0
|
||||
assert analysis["matched_historical_events"][0]["news_id"] == historical_news_id
|
||||
assert analysis["matched_historical_events"][0]["similarity"] == 1.0
|
||||
assert analysis["matched_historical_events"][0]["direction"] == "same"
|
||||
assert analysis["forecast"]["suggested_holding_days"] == 10
|
||||
assert analysis["forecast"]["returns"][-1]["expected_return"] == pytest.approx(0.1)
|
||||
assert "calculation_code" not in analysis
|
||||
|
|
@ -630,7 +631,7 @@ def test_historical_market_sample_rejects_look_ahead_and_incorrect_adjusted_retu
|
|||
AutoFinMarketSample.model_validate(sample)
|
||||
|
||||
|
||||
def test_market_calculation_clamps_and_reverses_negative_similarity():
|
||||
def test_market_calculation_equal_weights_and_reverses_opposite_direction_event():
|
||||
item = AutoFinEtfSelection.model_validate(
|
||||
{
|
||||
"etf_code": "518880.SH",
|
||||
|
|
@ -667,16 +668,45 @@ def test_market_calculation_clamps_and_reverses_negative_similarity():
|
|||
},
|
||||
],
|
||||
},
|
||||
{
|
||||
"reason": "黄金价格方向相同",
|
||||
"news_id": "20260602100000_efgh",
|
||||
"source_path": "daily/2026-06-02/auto_fin_news_data.jsonl",
|
||||
"event_time": "2026-06-02T10:00:00",
|
||||
"event_title": "黄金价格上涨",
|
||||
"event_content": "黄金价格出现明显上涨。",
|
||||
"market_entry": {
|
||||
"entry_time": "2026-06-02T15:00:00",
|
||||
"trade_date": "2026-06-02",
|
||||
"price_type": "close",
|
||||
"raw_price": 1.0,
|
||||
"adj_factor": 1.0,
|
||||
},
|
||||
"future_returns": [
|
||||
{
|
||||
"horizon": 1,
|
||||
"trade_date": "2026-06-03",
|
||||
"raw_close": 1.3,
|
||||
"adj_factor": 1.0,
|
||||
"cumulative_return": 0.3,
|
||||
},
|
||||
],
|
||||
},
|
||||
],
|
||||
},
|
||||
)
|
||||
selection = AutoFinMarketSelection.model_validate(
|
||||
{
|
||||
"matched_historical_events": [
|
||||
"same_direction_events": [
|
||||
{
|
||||
"reason": "机制和价格方向均相同",
|
||||
"news_id": "20260602100000_efgh",
|
||||
},
|
||||
],
|
||||
"opposite_direction_events": [
|
||||
{
|
||||
"reason": "机制可比但价格方向相反",
|
||||
"news_id": "20260601100000_abcd",
|
||||
"similarity": -2.0,
|
||||
},
|
||||
],
|
||||
},
|
||||
|
|
@ -684,11 +714,23 @@ def test_market_calculation_clamps_and_reverses_negative_similarity():
|
|||
|
||||
analysis = AutoFinMarketStep._calculate_analysis(item, history, selection)
|
||||
|
||||
assert analysis.matched_historical_events[0].similarity == -1.0
|
||||
assert analysis.matched_historical_events[0].weight == 1.0
|
||||
assert analysis.forecast.returns[0].expected_return == pytest.approx(-0.1)
|
||||
assert analysis.forecast.suggested_holding_days is None
|
||||
assert "加权预期收益没有正值" in analysis.limitations
|
||||
assert [match.direction for match in analysis.matched_historical_events] == ["same", "opposite"]
|
||||
assert [match.weight for match in analysis.matched_historical_events] == [0.5, 0.5]
|
||||
assert analysis.forecast.returns[0].expected_return == pytest.approx(0.1)
|
||||
assert analysis.forecast.suggested_holding_days == 1
|
||||
assert "相似历史样本的收益方向存在分歧" in analysis.limitations
|
||||
|
||||
|
||||
def test_market_selection_rejects_news_repeated_across_direction_groups():
|
||||
duplicate = {"reason": "方向判断", "news_id": "20260601100000_abcd"}
|
||||
|
||||
with pytest.raises(ValueError, match="news IDs must be unique"):
|
||||
AutoFinMarketSelection.model_validate(
|
||||
{
|
||||
"same_direction_events": [duplicate],
|
||||
"opposite_direction_events": [duplicate],
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
|
|||
|
|
@ -3,6 +3,7 @@
|
|||
# pylint: disable=missing-function-docstring,protected-access
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
|
|
@ -51,8 +52,9 @@ class _Handler:
|
|||
|
||||
|
||||
class _WebSocket:
|
||||
def __init__(self, messages=()):
|
||||
def __init__(self, messages=(), wait_when_empty=True):
|
||||
self.messages = list(messages)
|
||||
self.wait_when_empty = wait_when_empty
|
||||
self.closed = asyncio.Event()
|
||||
|
||||
async def __aenter__(self):
|
||||
|
|
@ -67,6 +69,8 @@ class _WebSocket:
|
|||
async def __anext__(self):
|
||||
if self.messages:
|
||||
return self.messages.pop(0)
|
||||
if not self.wait_when_empty:
|
||||
raise StopAsyncIteration
|
||||
await self.closed.wait()
|
||||
raise StopAsyncIteration
|
||||
|
||||
|
|
@ -227,7 +231,57 @@ async def test_stream_client_closes_when_background_stop_is_set(monkeypatch):
|
|||
|
||||
@pytest.mark.asyncio
|
||||
async def test_stream_client_restarts_after_server_disconnect(monkeypatch):
|
||||
websocket = _WebSocket(["{}"])
|
||||
websocket = _WebSocket(
|
||||
[
|
||||
json.dumps(
|
||||
{
|
||||
"type": "SYSTEM",
|
||||
"headers": {"topic": "disconnect"},
|
||||
"data": json.dumps({"reason": "connection is expired"}),
|
||||
},
|
||||
),
|
||||
],
|
||||
)
|
||||
monkeypatch.setattr("websockets.connect", lambda _uri: websocket)
|
||||
with pytest.raises(ConnectionError, match="WebSocket closed"):
|
||||
await DingTalkWaitStep._run_client(_StreamClient("disconnect"), asyncio.Event())
|
||||
reason = await DingTalkWaitStep._run_client(_StreamClient("disconnect"), asyncio.Event())
|
||||
assert reason == "connection is expired"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_stream_client_raises_when_websocket_closes_unexpectedly(monkeypatch):
|
||||
websocket = _WebSocket(wait_when_empty=False)
|
||||
monkeypatch.setattr("websockets.connect", lambda _uri: websocket)
|
||||
with pytest.raises(ConnectionError, match="closed unexpectedly"):
|
||||
await DingTalkWaitStep._run_client(_StreamClient(), asyncio.Event())
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_stream_client_reconnects_after_server_request(monkeypatch, tmp_path):
|
||||
app_context = ApplicationContext(workspace_dir=str(tmp_path))
|
||||
step = DingTalkWaitStep(app_context=app_context)
|
||||
step.logger = MagicMock()
|
||||
stop_event = asyncio.Event()
|
||||
calls = 0
|
||||
|
||||
async def run_client(_client, _stop_event):
|
||||
nonlocal calls
|
||||
calls += 1
|
||||
if calls == 1:
|
||||
return "connection is expired"
|
||||
stop_event.set()
|
||||
return None
|
||||
|
||||
async def timeout(awaitable, *, timeout):
|
||||
del timeout
|
||||
awaitable.close()
|
||||
raise asyncio.TimeoutError
|
||||
|
||||
monkeypatch.setattr(step, "_run_client", run_client)
|
||||
monkeypatch.setattr(asyncio, "wait_for", timeout)
|
||||
|
||||
await step._run_with_reconnect(_StreamClient(), stop_event)
|
||||
|
||||
assert calls == 2
|
||||
step.logger.info.assert_called_once_with(
|
||||
"[DingTalkWaitStep] DingTalk server requested reconnect reason='connection is expired'; reconnecting in 1.0s",
|
||||
)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue