diff --git a/reme/schema/__init__.py b/reme/schema/__init__.py index 6a25163b..2498137c 100644 --- a/reme/schema/__init__.py +++ b/reme/schema/__init__.py @@ -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", diff --git a/reme/schema/auto_fin.py b/reme/schema/auto_fin.py index 3c55eb93..e1647294 100644 --- a/reme/schema/auto_fin.py +++ b/reme/schema/auto_fin.py @@ -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) diff --git a/reme/steps/cookbook/auto_fin/market.py b/reme/steps/cookbook/auto_fin/market.py index 2e4cd046..74c4ae4d 100644 --- a/reme/steps/cookbook/auto_fin/market.py +++ b/reme/steps/cookbook/auto_fin/market.py @@ -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, diff --git a/reme/steps/cookbook/auto_fin/market.yaml b/reme/steps/cookbook/auto_fin/market.yaml index a4b6308a..7e3d3edd 100644 --- a/reme/steps/cookbook/auto_fin/market.yaml +++ b/reme/steps/cookbook/auto_fin/market.yaml @@ -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 正文。 diff --git a/reme/steps/cookbook/dingtalk/wait.py b/reme/steps/cookbook/dingtalk/wait.py index afd370d3..618cd534 100644 --- a/reme/steps/cookbook/dingtalk/wait.py +++ b/reme/steps/cookbook/dingtalk/wait.py @@ -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") diff --git a/tests/unit/test_auto_fin.py b/tests/unit/test_auto_fin.py index 5cdf6709..1da5784d 100644 --- a/tests/unit/test_auto_fin.py +++ b/tests/unit/test_auto_fin.py @@ -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 diff --git a/tests/unit/test_dingtalk_wait.py b/tests/unit/test_dingtalk_wait.py index a3836361..09768971 100644 --- a/tests/unit/test_dingtalk_wait.py +++ b/tests/unit/test_dingtalk_wait.py @@ -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", + )