From 3c2a9b0a7100c6f5d46d299451d7f5a6c4cfbc59 Mon Sep 17 00:00:00 2001 From: "jinli.yl" Date: Fri, 18 Sep 2026 15:59:37 +0800 Subject: [PATCH] fix(auto-fin): parse topic IDs from fenced JSON replies --- plugins/auto-fin/src/reme_auto_fin/topic.py | 50 ++++++++++-- plugins/auto-fin/src/reme_auto_fin/topic.yaml | 8 ++ plugins/auto-fin/tests/test_auto_fin.py | 76 +++++++++++++++++-- 3 files changed, 122 insertions(+), 12 deletions(-) diff --git a/plugins/auto-fin/src/reme_auto_fin/topic.py b/plugins/auto-fin/src/reme_auto_fin/topic.py index d5c5e680..df9c3c2c 100644 --- a/plugins/auto-fin/src/reme_auto_fin/topic.py +++ b/plugins/auto-fin/src/reme_auto_fin/topic.py @@ -3,15 +3,56 @@ from __future__ import annotations import json +import re +from time import perf_counter -from .base import AutoFinStep -from .schema import AutoFinTopicOutput +from .base import AGENT_INPUT_LOG_LIMIT, AGENT_OUTPUT_LOG_LIMIT, AutoFinStep class AutoFinTopicStep(AutoFinStep): """Filter current news in bounded Agent batches without writing files.""" + @staticmethod + def _parse_news_ids(value: object) -> list[str]: + if not isinstance(value, str): + raise ValueError("Auto Fin topic Agent returned no text") + match = re.search(r"```json\s*(.*?)```", value, re.IGNORECASE | re.DOTALL) + if match is None: + raise ValueError("Auto Fin topic Agent returned no JSON code block") + try: + ids = json.loads(match.group(1).strip()) + except json.JSONDecodeError as exc: + raise ValueError("Auto Fin topic Agent returned invalid JSON") from exc + if not isinstance(ids, list) or any(not isinstance(item, str) for item in ids): + raise ValueError("Auto Fin topic Agent must return a JSON array of strings") + return ids + + async def _select_news_ids(self, prompt: str) -> list[str]: + """Request and validate plain-text topic IDs, retrying one malformed reply.""" + if self.agent_wrapper is None: + raise RuntimeError("Auto Fin analysis requires an agent_wrapper") + self.logger.info( + f"[{self.name}] agent input prompt=topic_user query={self._preview(prompt, AGENT_INPUT_LOG_LIMIT)}", + ) + for attempt in range(2): + started_at = perf_counter() + result = await self.agent_wrapper.reply(prompt) + try: + ids = self._parse_news_ids(result.get("result") if isinstance(result, dict) else None) + except ValueError as exc: + if attempt: + raise ValueError(f"Auto Fin topic Agent returned invalid news IDs: {exc}") from exc + self.logger.warning(f"[{self.name}] invalid topic JSON; retrying once: {exc}") + continue + self.logger.info( + f"[{self.name}] agent output prompt=topic_user elapsed={perf_counter() - started_at:.2f}s " + f"output={self._preview(ids, AGENT_OUTPUT_LOG_LIMIT)}", + ) + return ids + raise RuntimeError("Auto Fin topic Agent produced no response") + async def execute(self): + """Select relevant news from each batch for the current invocation.""" assert self.context is not None news = list(self._required("auto_fin_news")) topics = list(self._required("auto_fin_topics")) @@ -23,14 +64,13 @@ class AutoFinTopicStep(AutoFinStep): batch = [ {**row, "content": str(row.get("content") or "")[:1000]} for row in news[start : start + batch_size] ] - output = await self._reply( + prompt = self.prompt_format( "topic_user", - AutoFinTopicOutput, topics=json.dumps(topics, ensure_ascii=False), news=json.dumps(batch, ensure_ascii=False), window_hours=formatted_hours, ) - selected.update(output.news_ids) + selected.update(await self._select_news_ids(prompt)) relevant = [row for row in news if row["news_id"] in selected] self.context["auto_fin_selected_news"] = relevant self.context.response.metadata["relevant_news_count"] = len(relevant) diff --git a/plugins/auto-fin/src/reme_auto_fin/topic.yaml b/plugins/auto-fin/src/reme_auto_fin/topic.yaml index 638e9a7d..f0035e6e 100644 --- a/plugins/auto-fin/src/reme_auto_fin/topic.yaml +++ b/plugins/auto-fin/src/reme_auto_fin/topic.yaml @@ -1,6 +1,14 @@ topic_user: | 从下面最近{window_hours}小时的财联社新闻中,选择与至少一个 topics 存在真实、可解释关系的新闻。 只返回输入中真实存在的 news_id;仅出现关键词但没有实质关系的新闻不要选择。 + 在 ```json 代码块中返回 JSON 字符串数组。没有相关新闻时返回 []。 + 输出格式示例: + ```json + [ + "123", + "456" + ] + ``` topics:{topics} 新闻:{news} diff --git a/plugins/auto-fin/tests/test_auto_fin.py b/plugins/auto-fin/tests/test_auto_fin.py index a273d68d..4336c818 100644 --- a/plugins/auto-fin/tests/test_auto_fin.py +++ b/plugins/auto-fin/tests/test_auto_fin.py @@ -2,6 +2,7 @@ # pylint: disable=missing-function-docstring,protected-access +import json from datetime import datetime from pathlib import Path from zoneinfo import ZoneInfo @@ -25,7 +26,12 @@ PLUGIN_MANIFEST = yaml.safe_load( def _row(news_id: int, value: datetime, title: str = "新闻", content: str = "正文") -> dict: - return {"id": news_id, "ctime": int(value.timestamp()), "title": title, "content": content} + return { + "id": news_id, + "ctime": int(value.timestamp()), + "title": title, + "content": content, + } def test_atomic_write_preserves_existing_file_on_failure(tmp_path: Path, monkeypatch): @@ -101,7 +107,7 @@ class _TopicAgent(BaseAgentWrapper): async def reply(self, inputs, **kwargs): self.calls.append((str(inputs), kwargs)) - return {"structured_output": AutoFinTopicOutput(news_ids=self.news_ids)} + return {"result": f"筛选结果:\n```json\n{json.dumps(self.news_ids)}\n```\n以上是相关 ID。"} @pytest.mark.asyncio @@ -110,8 +116,18 @@ async def test_topic_step_keeps_real_ids_in_memory_only(tmp_path: Path): agent = _TopicAgent(["2", "missing", "2"], app_context=app_context) context = RuntimeContext( auto_fin_news=[ - {"news_id": "1", "event_time": "2026-08-10T08:00:00+08:00", "title": "甲", "content": "甲"}, - {"news_id": "2", "event_time": "2026-08-10T09:00:00+08:00", "title": "乙", "content": "乙"}, + { + "news_id": "1", + "event_time": "2026-08-10T08:00:00+08:00", + "title": "甲", + "content": "甲", + }, + { + "news_id": "2", + "event_time": "2026-08-10T09:00:00+08:00", + "title": "乙", + "content": "乙", + }, ], auto_fin_topics=["黄金"], ) @@ -119,7 +135,8 @@ async def test_topic_step_keeps_real_ids_in_memory_only(tmp_path: Path): response = await AutoFinTopicStep(app_context=app_context, agent_wrapper=agent)(context) assert [row["news_id"] for row in context["auto_fin_selected_news"]] == ["2"] - assert agent.calls[0][1] == {"output_schema": AutoFinTopicOutput} + assert agent.calls[0][1] == {} + assert '```json\n[\n"123",\n"456"\n]\n```' in agent.calls[0][0] assert response.metadata["relevant_news_count"] == 1 assert not list(tmp_path.rglob("*.*")) @@ -130,7 +147,12 @@ async def test_topic_step_marks_empty_selection_as_successful_skip(tmp_path: Pat agent = _TopicAgent([], app_context=app_context) context = RuntimeContext( auto_fin_news=[ - {"news_id": "1", "event_time": "2026-08-10T08:00:00+08:00", "title": "甲", "content": "甲"}, + { + "news_id": "1", + "event_time": "2026-08-10T08:00:00+08:00", + "title": "甲", + "content": "甲", + }, ], auto_fin_topics=["黄金"], auto_fin_window_hours=12, @@ -148,6 +170,44 @@ async def test_topic_step_marks_empty_selection_as_successful_skip(tmp_path: Pat assert not list(tmp_path.rglob("*.md")) +@pytest.mark.asyncio +async def test_topic_step_retries_invalid_json_once(tmp_path: Path): + class RetryAgent(_TopicAgent): + """Return one malformed response before a fenced JSON array.""" + + async def reply(self, inputs, **kwargs): + self.calls.append((str(inputs), kwargs)) + return {"result": ('{"new_ids": "[\\"1\\"]"}' if len(self.calls) == 1 else '```json\n["1"]\n```')} + + app_context = ApplicationContext(workspace_dir=str(tmp_path), timezone="Asia/Shanghai") + agent = RetryAgent([], app_context=app_context) + context = RuntimeContext( + auto_fin_news=[ + { + "news_id": "1", + "event_time": "2026-08-10T08:00:00+08:00", + "title": "甲", + "content": "甲", + }, + ], + auto_fin_topics=["黄金"], + ) + + await AutoFinTopicStep(app_context=app_context, agent_wrapper=agent)(context) + + assert len(agent.calls) == 2 + assert [row["news_id"] for row in context["auto_fin_selected_news"]] == ["1"] + + +@pytest.mark.parametrize( + "value", + ['{"new_ids": ["1"]}', "```json\n[1]\n```", "not json", ""], +) +def test_topic_step_rejects_non_array_or_non_string_ids(value: str): + with pytest.raises(ValueError): + AutoFinTopicStep._parse_news_ids(value) + + class _ResearchAgent(BaseAgentWrapper): def __init__(self, **kwargs): super().__init__(**kwargs) @@ -171,7 +231,9 @@ class _ResearchAgent(BaseAgentWrapper): @pytest.mark.asyncio -async def test_merge_writes_only_final_report_and_validates_historical_links(tmp_path: Path): +async def test_merge_writes_only_final_report_and_validates_historical_links( + tmp_path: Path, +): historical = tmp_path / "daily" / "2026-08-01" / "auto_fin.md" historical.parent.mkdir(parents=True) historical.write_text("# 历史黄金观察\n", encoding="utf-8")