feat: 日报事件多新闻源入库(news_event.sources 列)

- ExtractedEvent 增加 sources 字段(主源居首去重,旧产物兜底 [source_id])
- extract_event(_async) 增加 sources 参数;run_event_extraction 读取
  deduped uniques 的 sources 透传
- reporter _load_events_from_dir 读取 sources → EventRow.sources
- news_event 建表加 sources TEXT 列;init_schema 幂等迁移旧表(捕获
  Duplicate column);save_report INSERT 写 JSON 多源
- 新增测试:ExtractedEvent/EventRow.sources 兜底与保序、_load_events_from_dir 多源
This commit is contained in:
2026-08-12 12:56:04 +08:00
parent 182f788857
commit 794ce6c853
10 changed files with 129 additions and 7 deletions
+6
View File
@@ -188,10 +188,12 @@ def extract_event(
*, *,
template: PromptTemplate | None = None, template: PromptTemplate | None = None,
max_attempts: int | None = None, max_attempts: int | None = None,
sources: list[str] | None = None,
) -> ExtractedEvent: ) -> ExtractedEvent:
"""同步抽取单篇文章的事件(带重试)。 """同步抽取单篇文章的事件(带重试)。
max_attempts 为 None 时使用 config.max_attempts(来自 YAML/环境变量配置)。 max_attempts 为 None 时使用 config.max_attempts(来自 YAML/环境变量配置)。
sources 为该新闻全部来源(来自去重层多源记录);None 时兜底 [article.source_id]。
""" """
tpl = template or PromptTemplate() tpl = template or PromptTemplate()
prompt = tpl.render(article) prompt = tpl.render(article)
@@ -208,6 +210,7 @@ def extract_event(
url_hash=article.url_hash, url_hash=article.url_hash,
title=article.title, title=article.title,
publish_time=article.publish_time, publish_time=article.publish_time,
sources=sources or [article.source_id],
event=event, event=event,
provider=config.provider, provider=config.provider,
model=config.model, model=config.model,
@@ -246,11 +249,13 @@ async def extract_event_async(
*, *,
template: PromptTemplate | None = None, template: PromptTemplate | None = None,
max_attempts: int | None = None, max_attempts: int | None = None,
sources: list[str] | None = None,
semaphore: asyncio.Semaphore | None = None, semaphore: asyncio.Semaphore | None = None,
) -> ExtractedEvent: ) -> ExtractedEvent:
"""异步抽取(批处理用),与同步版逻辑等价。 """异步抽取(批处理用),与同步版逻辑等价。
max_attempts 为 None 时使用 config.max_attempts。 max_attempts 为 None 时使用 config.max_attempts。
sources 为该新闻全部来源;None 时兜底 [article.source_id]。
""" """
tpl = template or PromptTemplate() tpl = template or PromptTemplate()
prompt = tpl.render(article) prompt = tpl.render(article)
@@ -268,6 +273,7 @@ async def extract_event_async(
url_hash=article.url_hash, url_hash=article.url_hash,
title=article.title, title=article.title,
publish_time=article.publish_time, publish_time=article.publish_time,
sources=sources or [article.source_id],
event=event, event=event,
provider=config.provider, provider=config.provider,
model=config.model, model=config.model,
+20 -1
View File
@@ -130,7 +130,12 @@ class EventExtraction(BaseModel):
class ExtractedEvent(BaseModel): class ExtractedEvent(BaseModel):
"""落盘格式:文章元数据 + LLM 抽取结果 + 调用元信息。""" """落盘格式:文章元数据 + LLM 抽取结果 + 调用元信息。
sources: 该唯一新闻的全部来源(主源 source_id 居首)。
来自去重层多源记录(M3 sources.json / uniques JSON 的 sources 字段),
旧产物无此字段时兜底为 [source_id]。
"""
# ---- 来源标识 ---- # ---- 来源标识 ----
source_id: str source_id: str
@@ -138,6 +143,10 @@ class ExtractedEvent(BaseModel):
url_hash: str url_hash: str
title: str title: str
publish_time: datetime | None = None publish_time: datetime | None = None
sources: list[str] = Field(
default_factory=list,
description="全部来源(主源居首,去重保序);旧产物无字段时兜底为 [source_id]",
)
# ---- 抽取结果 ---- # ---- 抽取结果 ----
event: EventExtraction event: EventExtraction
@@ -150,6 +159,16 @@ class ExtractedEvent(BaseModel):
prompt_tokens: int | None = None prompt_tokens: int | None = None
completion_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: def short_summary(self) -> str:
ev = self.event ev = self.event
codes = ",".join(ev.stock_codes) or "-" codes = ",".join(ev.stock_codes) or "-"
+14 -4
View File
@@ -11,7 +11,7 @@ import pymysql
from loguru import logger from loguru import logger
from .models import ReportData from .models import ReportData
from .schema import DDL_STATEMENTS from .schema import _ALTER_ADD_SOURCES, DDL_STATEMENTS
@dataclass(frozen=True) @dataclass(frozen=True)
@@ -60,10 +60,19 @@ def connect(cfg: DbConfig | None = None) -> pymysql.Connection:
def init_schema(conn: pymysql.Connection) -> None: 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: with conn.cursor() as cur:
for ddl in DDL_STATEMENTS: for ddl in DDL_STATEMENTS:
cur.execute(ddl) 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() conn.commit()
logger.info("news_report / news_event 建表完成") logger.info("news_report / news_event 建表完成")
@@ -112,8 +121,8 @@ def save_report(conn: pymysql.Connection, report: ReportData) -> int:
""" """
INSERT INTO news_event INSERT INTO news_event
(report_id, section, rank, importance, event_type, title, (report_id, section, rank, importance, event_type, title,
summary, sentiment, source, url) summary, sentiment, source, sources, url)
VALUES (%s, %s, %s, %s, %s, %s, %s, %s, %s, %s) VALUES (%s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s)
""", """,
( (
report_id, report_id,
@@ -125,6 +134,7 @@ def save_report(conn: pymysql.Connection, report: ReportData) -> int:
ev.summary, ev.summary,
ev.sentiment, ev.sentiment,
ev.source, ev.source,
json.dumps(ev.sources, ensure_ascii=False) if ev.sources else None,
ev.url, ev.url,
), ),
) )
+2 -1
View File
@@ -18,7 +18,8 @@ class EventRow(BaseModel):
title: str title: str
summary: str | None = None summary: str | None = None
sentiment: str | None = None # positive | negative | neutral | '' sentiment: str | None = None # positive | negative | neutral | ''
source: str | None = None source: str | None = None # 主源
sources: list[str] | None = None # 全部来源(主源居首),多源新闻记录
url: str | None = None url: str | None = None
+8 -1
View File
@@ -33,7 +33,8 @@ DDL_STATEMENTS: list[str] = [
title VARCHAR(512) NOT NULL COMMENT '标题', title VARCHAR(512) NOT NULL COMMENT '标题',
summary TEXT NULL COMMENT '摘要/正文', summary TEXT NULL COMMENT '摘要/正文',
sentiment VARCHAR(8) NULL COMMENT 'positive/negative/neutral', 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 '原文链接', url VARCHAR(512) NULL COMMENT '原文链接',
created_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP, created_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP,
KEY idx_report_section (report_id, section), 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='日报事件明细' ) 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"
)
+2
View File
@@ -102,6 +102,7 @@ def _load_events_from_dir(day_str: str) -> list[dict]:
"title": obj.get("title", ""), "title": obj.get("title", ""),
"url": obj.get("url", ""), "url": obj.get("url", ""),
"source_id": obj.get("source_id", ""), "source_id": obj.get("source_id", ""),
"sources": obj.get("sources") or [obj.get("source_id", "")],
"publish_time": obj.get("publish_time"), "publish_time": obj.get("publish_time"),
"event": ev, "event": ev,
}) })
@@ -1004,6 +1005,7 @@ def _build_report_data(news: dict, cninfo: dict, pipeline: dict,
summary=(ev.get("summary") or None), summary=(ev.get("summary") or None),
sentiment=ev.get("sentiment") or None, sentiment=ev.get("sentiment") or None,
source=e.get("source_id") or None, source=e.get("source_id") or None,
sources=e.get("sources") or None,
url=e.get("url") or None, url=e.get("url") or None,
) )
) )
+17
View File
@@ -109,6 +109,21 @@ def _load_article(p: Path) -> Article | None:
return 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: async def _run(args: argparse.Namespace) -> int:
load_dotenv() # 读 .env 到 os.environ load_dotenv() # 读 .env 到 os.environ
# scene=event_extraction: 读取 configs/llm_models.yaml 场景 1 配置,未配置字段回退 .env # scene=event_extraction: 读取 configs/llm_models.yaml 场景 1 配置,未配置字段回退 .env
@@ -173,6 +188,7 @@ async def _run(args: argparse.Namespace) -> int:
article = _load_article(fp) article = _load_article(fp)
if article is None: if article is None:
return fp, None, "无法解析输入" return fp, None, "无法解析输入"
sources = _load_sources(fp) # 多源记录(去重层),无则 None
try: try:
event = await extract_event_async( event = await extract_event_async(
client=client, client=client,
@@ -180,6 +196,7 @@ async def _run(args: argparse.Namespace) -> int:
article=article, article=article,
template=template, template=template,
max_attempts=args.max_attempts, max_attempts=args.max_attempts,
sources=sources,
semaphore=semaphore, semaphore=semaphore,
) )
return fp, event, None return fp, event, None
+25
View File
@@ -121,6 +121,31 @@ def test_event_types_constant_includes_common() -> None:
assert must in EVENT_TYPES 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 提取与解析 # JSON 提取与解析
# --------------------------------------------------------------------------- # # --------------------------------------------------------------------------- #
+28
View File
@@ -255,3 +255,31 @@ class TestCollectNewsEventsLookback:
assert "无时间戳" in titles assert "无时间戳" in titles
assert "窗口外旧闻" not in titles assert "窗口外旧闻" not 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"]
+7
View File
@@ -20,9 +20,16 @@ class TestEventRow:
ev = EventRow( ev = EventRow(
section="intl", rank=2, importance=4, event_type="地缘政治", section="intl", rank=2, importance=4, event_type="地缘政治",
title="t", summary="s", sentiment="negative", source="ForexLive", title="t", summary="s", sentiment="negative", source="ForexLive",
sources=["ForexLive", "新浪财经"],
url="https://x.com/1", url="https://x.com/1",
) )
assert ev.sentiment == "negative" 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: def test_missing_title_raises(self) -> None:
with pytest.raises(ValidationError): with pytest.raises(ValidationError):