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:
jinliyl 2026-07-27 11:52:23 +08:00 • committed by GitHub
parent 0522135791
commit 2f79977df0
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
7 changed files with 200 additions and 71 deletions

View file

@ -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",

View file

@ -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)

View file

@ -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,

View file

@ -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 正文。

View file

@ -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")

View file

@ -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

View file

@ -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",
)