mirror of
https://github.com/agentscope-ai/ReMe.git
synced 2026-10-01 02:04:21 +00:00
fix(auto-fin): parse topic IDs from fenced JSON replies
This commit is contained in:
parent
5231f3970c
commit
3c2a9b0a71
3 changed files with 122 additions and 12 deletions
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -1,6 +1,14 @@
|
|||
topic_user: |
|
||||
从下面最近{window_hours}小时的财联社新闻中,选择与至少一个 topics 存在真实、可解释关系的新闻。
|
||||
只返回输入中真实存在的 news_id;仅出现关键词但没有实质关系的新闻不要选择。
|
||||
在 ```json 代码块中返回 JSON 字符串数组。没有相关新闻时返回 []。
|
||||
输出格式示例:
|
||||
```json
|
||||
[
|
||||
"123",
|
||||
"456"
|
||||
]
|
||||
```
|
||||
|
||||
topics:{topics}
|
||||
新闻:{news}
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue