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:
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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"
|
||||||
|
)
|
||||||
|
|||||||
@@ -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,
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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 提取与解析
|
||||||
# --------------------------------------------------------------------------- #
|
# --------------------------------------------------------------------------- #
|
||||||
|
|||||||
@@ -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"]
|
||||||
|
|||||||
@@ -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):
|
||||||
|
|||||||
Reference in New Issue
Block a user