diff --git a/llm/extractor.py b/llm/extractor.py index 3b09fba..adc2d12 100644 --- a/llm/extractor.py +++ b/llm/extractor.py @@ -188,10 +188,12 @@ def extract_event( *, template: PromptTemplate | None = None, max_attempts: int | None = None, + sources: list[str] | None = None, ) -> ExtractedEvent: """同步抽取单篇文章的事件(带重试)。 max_attempts 为 None 时使用 config.max_attempts(来自 YAML/环境变量配置)。 + sources 为该新闻全部来源(来自去重层多源记录);None 时兜底 [article.source_id]。 """ tpl = template or PromptTemplate() prompt = tpl.render(article) @@ -208,6 +210,7 @@ def extract_event( url_hash=article.url_hash, title=article.title, publish_time=article.publish_time, + sources=sources or [article.source_id], event=event, provider=config.provider, model=config.model, @@ -246,11 +249,13 @@ async def extract_event_async( *, template: PromptTemplate | None = None, max_attempts: int | None = None, + sources: list[str] | None = None, semaphore: asyncio.Semaphore | None = None, ) -> ExtractedEvent: """异步抽取(批处理用),与同步版逻辑等价。 max_attempts 为 None 时使用 config.max_attempts。 + sources 为该新闻全部来源;None 时兜底 [article.source_id]。 """ tpl = template or PromptTemplate() prompt = tpl.render(article) @@ -268,6 +273,7 @@ async def extract_event_async( url_hash=article.url_hash, title=article.title, publish_time=article.publish_time, + sources=sources or [article.source_id], event=event, provider=config.provider, model=config.model, diff --git a/llm/models.py b/llm/models.py index 60d16ed..0518287 100644 --- a/llm/models.py +++ b/llm/models.py @@ -130,7 +130,12 @@ class EventExtraction(BaseModel): class ExtractedEvent(BaseModel): - """落盘格式:文章元数据 + LLM 抽取结果 + 调用元信息。""" + """落盘格式:文章元数据 + LLM 抽取结果 + 调用元信息。 + + sources: 该唯一新闻的全部来源(主源 source_id 居首)。 + 来自去重层多源记录(M3 sources.json / uniques JSON 的 sources 字段), + 旧产物无此字段时兜底为 [source_id]。 + """ # ---- 来源标识 ---- source_id: str @@ -138,6 +143,10 @@ class ExtractedEvent(BaseModel): url_hash: str title: str publish_time: datetime | None = None + sources: list[str] = Field( + default_factory=list, + description="全部来源(主源居首,去重保序);旧产物无字段时兜底为 [source_id]", + ) # ---- 抽取结果 ---- event: EventExtraction @@ -150,6 +159,16 @@ class ExtractedEvent(BaseModel): prompt_tokens: int | None = None completion_tokens: int | None = None + @model_validator(mode="after") + def _ensure_sources(self) -> Self: + """保证 sources 非空、去重且以主源 source_id 开头。""" + seen: list[str] = [] + for s in [self.source_id, *self.sources]: + if s and s not in seen: + seen.append(s) + self.sources = seen + return self + def short_summary(self) -> str: ev = self.event codes = ",".join(ev.stock_codes) or "-" diff --git a/report_db/db.py b/report_db/db.py index e1993ff..1787fd3 100644 --- a/report_db/db.py +++ b/report_db/db.py @@ -11,7 +11,7 @@ import pymysql from loguru import logger from .models import ReportData -from .schema import DDL_STATEMENTS +from .schema import _ALTER_ADD_SOURCES, DDL_STATEMENTS @dataclass(frozen=True) @@ -60,10 +60,19 @@ def connect(cfg: DbConfig | None = None) -> pymysql.Connection: def init_schema(conn: pymysql.Connection) -> None: - """建表(CREATE TABLE IF NOT EXISTS ×2),幂等。""" + """建表(CREATE TABLE IF NOT EXISTS ×2)+ 旧表 sources 列迁移,幂等。""" with conn.cursor() as cur: for ddl in DDL_STATEMENTS: cur.execute(ddl) + # 旧库兼容:news_event 已存在但缺 sources 列时补列 + try: + cur.execute(_ALTER_ADD_SOURCES) + logger.info("news_event 迁移:新增 sources 多源列") + except pymysql.err.OperationalError as e: + if e.args and "Duplicate column" in str(e.args[0]): + logger.debug("news_event.sources 列已存在,跳过迁移") + else: + raise conn.commit() logger.info("news_report / news_event 建表完成") @@ -112,8 +121,8 @@ def save_report(conn: pymysql.Connection, report: ReportData) -> int: """ INSERT INTO news_event (report_id, section, rank, importance, event_type, title, - summary, sentiment, source, url) - VALUES (%s, %s, %s, %s, %s, %s, %s, %s, %s, %s) + summary, sentiment, source, sources, url) + VALUES (%s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s) """, ( report_id, @@ -125,6 +134,7 @@ def save_report(conn: pymysql.Connection, report: ReportData) -> int: ev.summary, ev.sentiment, ev.source, + json.dumps(ev.sources, ensure_ascii=False) if ev.sources else None, ev.url, ), ) diff --git a/report_db/models.py b/report_db/models.py index 63b2e81..efea1a4 100644 --- a/report_db/models.py +++ b/report_db/models.py @@ -18,7 +18,8 @@ class EventRow(BaseModel): title: str summary: str | None = None sentiment: str | None = None # positive | negative | neutral | '' - source: str | None = None + source: str | None = None # 主源 + sources: list[str] | None = None # 全部来源(主源居首),多源新闻记录 url: str | None = None diff --git a/report_db/schema.py b/report_db/schema.py index a6b951a..6c5fb59 100644 --- a/report_db/schema.py +++ b/report_db/schema.py @@ -33,7 +33,8 @@ DDL_STATEMENTS: list[str] = [ title VARCHAR(512) NOT NULL COMMENT '标题', summary TEXT NULL COMMENT '摘要/正文', sentiment VARCHAR(8) NULL COMMENT 'positive/negative/neutral', - source VARCHAR(64) NULL COMMENT '来源', + source VARCHAR(64) NULL COMMENT '主来源', + sources TEXT NULL COMMENT '全部来源 JSON 数组(主源居首),多源新闻记录', url VARCHAR(512) NULL COMMENT '原文链接', created_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP, KEY idx_report_section (report_id, section), @@ -41,3 +42,9 @@ DDL_STATEMENTS: list[str] = [ ) ENGINE=InnoDB DEFAULT CHARSET=utf8mb4 COLLATE=utf8mb4_unicode_ci COMMENT='日报事件明细' """, ] + +# 旧库迁移:为已存在的 news_event 表补充 sources 列(MySQL 无 ADD COLUMN IF NOT EXISTS) +_ALTER_ADD_SOURCES = ( + "ALTER TABLE news_event ADD COLUMN sources TEXT NULL " + "COMMENT '全部来源 JSON 数组(主源居首),多源新闻记录' AFTER source" +) diff --git a/scheduler/reporter.py b/scheduler/reporter.py index 8957f0a..e9f9498 100644 --- a/scheduler/reporter.py +++ b/scheduler/reporter.py @@ -102,6 +102,7 @@ def _load_events_from_dir(day_str: str) -> list[dict]: "title": obj.get("title", ""), "url": obj.get("url", ""), "source_id": obj.get("source_id", ""), + "sources": obj.get("sources") or [obj.get("source_id", "")], "publish_time": obj.get("publish_time"), "event": ev, }) @@ -1004,6 +1005,7 @@ def _build_report_data(news: dict, cninfo: dict, pipeline: dict, summary=(ev.get("summary") or None), sentiment=ev.get("sentiment") or None, source=e.get("source_id") or None, + sources=e.get("sources") or None, url=e.get("url") or None, ) ) diff --git a/scripts/run_event_extraction.py b/scripts/run_event_extraction.py index 5b717ac..c791ea7 100644 --- a/scripts/run_event_extraction.py +++ b/scripts/run_event_extraction.py @@ -109,6 +109,21 @@ def _load_article(p: Path) -> Article | None: return None +def _load_sources(p: Path) -> list[str] | None: + """从输入 JSON 读取 sources 多源字段(deduped uniques 才有);无则返回 None。 + + None 表示无多源记录,由 extract_event 兜底为 [source_id]。 + """ + try: + data = json.loads(p.read_text(encoding="utf-8")) + srcs = data.get("sources") + if isinstance(srcs, list) and srcs: + return [s for s in srcs if s] + except (json.JSONDecodeError, OSError): + pass + return None + + async def _run(args: argparse.Namespace) -> int: load_dotenv() # 读 .env 到 os.environ # scene=event_extraction: 读取 configs/llm_models.yaml 场景 1 配置,未配置字段回退 .env @@ -173,6 +188,7 @@ async def _run(args: argparse.Namespace) -> int: article = _load_article(fp) if article is None: return fp, None, "无法解析输入" + sources = _load_sources(fp) # 多源记录(去重层),无则 None try: event = await extract_event_async( client=client, @@ -180,6 +196,7 @@ async def _run(args: argparse.Namespace) -> int: article=article, template=template, max_attempts=args.max_attempts, + sources=sources, semaphore=semaphore, ) return fp, event, None diff --git a/tests/test_llm.py b/tests/test_llm.py index e128ffa..dd90824 100644 --- a/tests/test_llm.py +++ b/tests/test_llm.py @@ -121,6 +121,31 @@ def test_event_types_constant_includes_common() -> None: assert must in EVENT_TYPES +# --------------------------------------------------------------------------- # +# ExtractedEvent.sources 多源字段 +# --------------------------------------------------------------------------- # + +def test_extracted_event_sources_defaults_to_main_source() -> None: + """未提供 sources 时兜底为 [source_id](兼容旧产物)。""" + ev = ExtractedEvent( + source_id="cls", url="https://x/1", url_hash="h1", title="t", + event=EventExtraction(sentiment="positive", importance=3, event_type="重大合同"), + provider="deepseek", model="m", + ) + assert ev.sources == ["cls"] + + +def test_extracted_event_sources_keeps_main_first_and_dedup() -> None: + """sources 保主源居首、去重保序。""" + ev = ExtractedEvent( + source_id="cls", url="https://x/1", url_hash="h1", title="t", + sources=["sina", "cls", "eastmoney", "sina"], + event=EventExtraction(sentiment="neutral", importance=2, event_type="其他"), + provider="deepseek", model="m", + ) + assert ev.sources == ["cls", "sina", "eastmoney"] + + # --------------------------------------------------------------------------- # # JSON 提取与解析 # --------------------------------------------------------------------------- # diff --git a/tests/test_report_builder.py b/tests/test_report_builder.py index 528f2a5..1916a41 100644 --- a/tests/test_report_builder.py +++ b/tests/test_report_builder.py @@ -255,3 +255,31 @@ class TestCollectNewsEventsLookback: assert "无时间戳" in titles assert "窗口外旧闻" not in titles assert "公告排除" not in titles + + +class TestLoadEventsSources: + """_load_events_from_dir 读取多源(sources)字段。""" + + def test_loads_sources_with_fallback(self, monkeypatch, tmp_path) -> None: + import json as _json + + import scheduler.reporter as rep + + ev_dir = tmp_path / "data" / "events" / "20260812" + ev_dir.mkdir(parents=True) + (ev_dir / "aaa.json").write_text(_json.dumps({ + "source_id": "cls", "sources": ["cls", "sina", "eastmoney"], + "title": "多源新闻", "url": "https://x/1", "event": {}, + }, ensure_ascii=False), encoding="utf-8") + (ev_dir / "bbb.json").write_text(_json.dumps({ + "source_id": "cls", "title": "旧产物无 sources", "url": "https://x/2", + "event": {}, + }, ensure_ascii=False), encoding="utf-8") + + monkeypatch.chdir(tmp_path) + events = rep._load_events_from_dir("20260812") + by_url = {e["url"]: e for e in events} + # 新产物:多源完整透传 + assert by_url["https://x/1"]["sources"] == ["cls", "sina", "eastmoney"] + # 旧产物:兜底 [source_id] + assert by_url["https://x/2"]["sources"] == ["cls"] diff --git a/tests/test_report_db.py b/tests/test_report_db.py index b13cda8..3c8898c 100644 --- a/tests/test_report_db.py +++ b/tests/test_report_db.py @@ -20,9 +20,16 @@ class TestEventRow: ev = EventRow( section="intl", rank=2, importance=4, event_type="地缘政治", title="t", summary="s", sentiment="negative", source="ForexLive", + sources=["ForexLive", "新浪财经"], url="https://x.com/1", ) assert ev.sentiment == "negative" + assert ev.sources == ["ForexLive", "新浪财经"] + + def test_sources_optional(self) -> None: + """sources 为可选项(旧数据无多源记录)。""" + ev = EventRow(section="news", rank=1, title="t") + assert ev.sources is None def test_missing_title_raises(self) -> None: with pytest.raises(ValidationError):