diff --git a/.gitignore b/.gitignore index 9d4e713..296bf1d 100644 --- a/.gitignore +++ b/.gitignore @@ -34,6 +34,9 @@ htmlcov/ .DS_Store Thumbs.db +# 工具本地配置(非项目文件) +reasonix.toml + # 项目敏感配置 .env .env.local diff --git a/configs/__init__.py b/configs/__init__.py new file mode 100644 index 0000000..a9b0e6d --- /dev/null +++ b/configs/__init__.py @@ -0,0 +1 @@ +"""configs 配置包:提供 configs/*.yaml 的加载能力。""" diff --git a/configs/llm_models.yaml b/configs/llm_models.yaml new file mode 100644 index 0000000..39fde39 --- /dev/null +++ b/configs/llm_models.yaml @@ -0,0 +1,133 @@ +# ============================================================================= +# configs/llm_models.yaml —— 大模型使用场景配置 +# +# 本文件集中配置本项目所有调用 AI 大模型的地方,每个场景可独立指定 +# provider(厂商/服务类型)与 model(模型名),互不影响、单独切换。 +# +# ── 配置优先级(从高到低)────────────────────────────────────────────────────── +# 1. 代码 / 命令行显式参数(如 --provider qwen --model qwen-plus) +# 2. 本文件 scenes.<场景>.xxx +# 3. 环境变量 / .env(LLM_PROVIDER、DEEPSEEK_MODEL 等,向后兼容) +# 4. 代码内置默认值 +# +# ── 规则 ───────────────────────────────────────────────────────────────────── +# · provider 取值: +# 对话大模型: deepseek | qwen(OpenAI 兼容 chat 接口) +# 向量嵌入模型: dashscope | local-bge(仅 embedding 场景) +# · model 留空 = 该场景不覆盖模型 → 回退 .env(如 DEEPSEEK_MODEL / QWEN_MODEL / +# LLM_MODEL);若全部缺失则直接报错,绝不静默使用内置默认模型。 +# · api_key_env / base_url_env 为可选字段,填写存放 API Key / 服务地址的 +# 环境变量名;API Key 一律放 .env,禁止写入本文件(安全规范)。 +# · 修改后无需重启常驻服务即可生效(每次调用重新读取;如需热更新缓存可重启)。 +# ============================================================================= + +# ---- 全局默认参数(各场景可覆盖;低于 .env,高于代码内置默认)---- +defaults: + timeout_sec: 60 + temperature: 0.1 + +scenes: + # ------------------------------------------------------------------------- # + # 场景 1: 投资事件抽取(M4 核心) + # ------------------------------------------------------------------------- # + # 用途:对每条唯一新闻调用大模型,抽取结构化投资事件 + # (stock_codes / company_names / industries / sentiment / importance / + # event_type / summary),输出严格 JSON,经 Pydantic 校验后落盘。 + # 调用方: llm/extractor.py、scripts/run_event_extraction.py、 + # scheduler/pipeline.py 的 llm 步骤。 + # 频率:每日数百~数千篇;建议异步并发(--concurrency,默认 3)。 + # 使用方式: + # uv run a-share events # 读本场景配置 + # uv run a-share events --provider qwen # CLI 临时覆盖 provider + # uv run a-share events --model qwen-max # CLI 临时覆盖模型 + # 对模型的要求: + # · 必须支持 OpenAI 兼容 chat/completions 接口; + # · 必须支持 JSON 结构化输出(response_format=json_object,硬性要求); + # · 上下文窗口 ≥ 16K tokens(单篇正文最多截断到 8000 字符); + # · 中文理解能力强,能区分 A 股事件类型(业绩预告/投资并购/宏观政策等 23 类); + # · 低 temperature(0.1)保证抽取稳定,避免字段抖动。 + # 建议模型: deepseek-v4-flash(生产实测) / deepseek-chat / qwen-plus / qwen-max + event_extraction: + provider: # 建议 deepseek | qwen;留空则回退 .env 的 LLM_PROVIDER + model: # 留空则回退 .env(DEEPSEEK_MODEL → LLM_MODEL) + api_key_env: # 例如: DEEPSEEK_API_KEY / QWEN_API_KEY / DASHSCOPE_API_KEY + base_url_env: # 例如: DEEPSEEK_BASE_URL / QWEN_BASE_URL + temperature: 0.1 + timeout_sec: 60 + max_attempts: 3 # 单篇解析失败的最大重试次数 + + # ------------------------------------------------------------------------- # + # 场景 2: 日报 AI 摘要 + # ------------------------------------------------------------------------- # + # 用途:每天汇总「新闻联播 + 过去 24h 高重要度新闻 + 近 N 日重要公告/调研」, + # 生成 500 字以内的日报要点摘要,渲染进 HTML 日报。 + # 调用方: scheduler/reporter.py 的 _generate_ai_summary → _llm_summarize。 + # 频率:每天 1 次(07:00 定时任务,与 pipeline 同批)。 + # 使用方式:无需手动触发,定时任务自动执行;失败自动降级(日报留空,不影响入库)。 + # 对模型的要求: + # · OpenAI 兼容 chat 接口(不需要 JSON 输出); + # · 输出长度 ≥ 1500 tokens(max_tokens=1500,输出超长会被截断并记 WARNING); + # · 中文摘要能力强、要点化输出稳定(每条一行,以 "- " 开头); + # · 上下文窗口 ≥ 8K tokens(素材按 3000 字符/块分块,多块先分段再合并); + # · temperature 0.3 左右,兼顾稳定与表达;网络失败按指数退避重试 3 次。 + # · 输出长度需求:分段摘要约 800 tokens、合并摘要约 1500 tokens(代码内置, + # 不在本文件配置),模型应能稳定输出 1500+ tokens 的中文要点。 + # 建议模型: deepseek-v4-flash(生产实测) / deepseek-chat / qwen-plus + daily_report: + provider: # 建议 deepseek | qwen;留空则回退 .env 的 LLM_PROVIDER + model: + api_key_env: + base_url_env: + temperature: 0.3 + timeout_sec: 60 + + # ------------------------------------------------------------------------- # + # 场景 3: 个股 AI 要点分析 + # ------------------------------------------------------------------------- # + # 用途:针对观察清单中的单只股票,汇总其近期公告/调研/新闻/互动问答, + # 生成 5-8 条要点分析,渲染进个股日报。 + # 调用方: scheduler/stock_reporter.py 的 _generate_ai_summary。 + # 频率:每交易日 1 次(07:30),对 watchlist 内每只股票各调用一次。 + # 使用方式:定时任务自动执行;单次失败不影响其他股票(该股显示"AI 摘要暂不可用")。 + # 对模型的要求: + # · OpenAI 兼容 chat 接口(不需要 JSON 输出); + # · 输出长度 ≥ 500 tokens(max_tokens=500); + # · 中文要点分析能力,输入素材最多 3500 字符(公告+调研+新闻+互动问答); + # · temperature 0.3 左右。 + # · 输出长度需求:约 500 tokens(代码内置,不在本文件配置)。 + # 建议模型: deepseek-v4-flash(生产实测) / deepseek-chat / qwen-plus + stock_report: + provider: # 建议 deepseek | qwen;留空则回退 .env 的 LLM_PROVIDER + model: + api_key_env: + base_url_env: + temperature: 0.3 + timeout_sec: 60 + + # ------------------------------------------------------------------------- # + # 场景 4: 文本向量化 Embedding(新闻/事件入库 + 语义检索) + # ------------------------------------------------------------------------- # + # 用途:把新闻正文/事件文本转为向量,写入 Qdrant 知识库;检索时对查询文本 + # 同样向量化后做相似度搜索。M5 入库、MCP 检索、pipeline 均复用本场景。 + # 调用方: embedding/remote.py(远程)、embedding/local.py(本地)、 + # embedding/factory.py、scripts/run_embedding.py、mcp_server/tools.py。 + # 使用方式: + # uv run a-share embed # 读本场景配置 + # uv run a-share embed --provider local-bge # 临时切换本地模型 + # 注意:这是「嵌入模型」而非对话大模型,二选一: + # · provider: dashscope → 阿里百炼 text-embedding-v3(远程,需 API key); + # · provider: local-bge → 本地 BGE-M3(离线,需 uv sync --extra + # local-embedding,首次加载约 2.3GB)。 + # 对模型的要求: + # · 输出固定维度向量(本项目默认 1024 维,DashScope 与本地 BGE-M3 兼容); + # · 中文语义匹配效果好,支持 batch 输入(单批上限 batch_limit 条); + # · 远程需 OpenAI 兼容 embeddings 接口。 + # 建议模型: text-embedding-v3(远程) / BAAI/bge-m3(本地) + embedding: + provider: # 建议 dashscope | local-bge;留空则回退 .env 的 EMBEDDING_PROVIDER + model: # 留空则回退 .env(DASHSCOPE_EMBEDDING_MODEL / LOCAL_EMBEDDING_MODEL) + api_key_env: # 例如: DASHSCOPE_EMBEDDING_API_KEY / DASHSCOPE_API_KEY + base_url_env: # 例如: DASHSCOPE_EMBEDDING_BASE_URL + timeout_sec: 60 + max_attempts: 3 # 单批请求失败重试次数 + batch_limit: 10 # 单批最大条数(百炼实测上限 10,勿调大) diff --git a/configs/loader.py b/configs/loader.py new file mode 100644 index 0000000..3179869 --- /dev/null +++ b/configs/loader.py @@ -0,0 +1,69 @@ +"""configs/ 目录下 YAML 配置加载器。 + +目前支持加载 configs/llm_models.yaml 的场景配置(scenes.)。 + +配置优先级(从高到低): + 1. 代码 / 命令行显式参数(如 --provider qwen --model qwen-plus) + 2. 本文件 YAML 场景配置(scenes.) + 3. 环境变量 / .env(LLM_PROVIDER、DEEPSEEK_MODEL 等,向后兼容) + 4. 代码内置默认值 + +说明:API Key 一律放 .env,本文件只保存环境变量名(api_key_env),禁止写密钥。 +""" + +from __future__ import annotations + +from functools import lru_cache +from pathlib import Path + +from loguru import logger + +DEFAULT_CONFIG_PATH = Path("configs/llm_models.yaml") + + +@lru_cache(maxsize=8) +def _load_yaml(path: Path) -> dict: + """读取 YAML 文件为 dict;文件缺失或解析失败返回空 dict(走兜底配置)。""" + if not path.is_file(): + logger.debug("配置文件不存在,使用内置/环境变量兜底: {}", path) + return {} + try: + import yaml + + data = yaml.safe_load(path.read_text(encoding="utf-8")) or {} + except Exception as e: # noqa: BLE001 - YAML 语法错误等 + logger.error("解析 {} 失败: {}", path, e) + return {} + return data if isinstance(data, dict) else {} + + +def load_scene_config(scene: str) -> dict: + """读取 llm_models.yaml 中 scenes. 的配置 dict。 + + 场景不存在或未配置时返回空 dict(调用方回退 .env / 内置默认)。 + scene 为空字符串同样返回空 dict。 + """ + if not scene: + return {} + data = _load_yaml(DEFAULT_CONFIG_PATH) + scenes = data.get("scenes") or {} + cfg = scenes.get(scene) + if cfg is None: + logger.debug("llm_models.yaml 未配置场景 {!r},使用 .env 兜底", scene) + return {} + if not isinstance(cfg, dict): + logger.warning("llm_models.yaml 场景 {!r} 应为 map,已忽略", scene) + return {} + return cfg + + +def load_defaults() -> dict: + """读取 llm_models.yaml 顶层 defaults(全局默认参数)。""" + data = _load_yaml(DEFAULT_CONFIG_PATH) + d = data.get("defaults") or {} + return d if isinstance(d, dict) else {} + + +def clear_cache() -> None: + """清空 YAML 缓存(测试或热更新配置时使用)。""" + _load_yaml.cache_clear() diff --git a/dedup/deduper.py b/dedup/deduper.py index 28ba822..1dfaf59 100644 --- a/dedup/deduper.py +++ b/dedup/deduper.py @@ -35,7 +35,10 @@ def _publish_date(article: Article) -> str | None: def article_to_fingerprint(article: Article) -> Fingerprint: - """构造 Fingerprint(用于 ingest 写入或对外只读)。""" + """构造 Fingerprint(用于 ingest 写入或对外只读)。 + + source_ids 初始为 [article.source_id],后续重复文章命中时由 ingest 合并。 + """ return Fingerprint( url_hash=article.url_hash, content_hash=content_hash(article.content), @@ -45,6 +48,7 @@ def article_to_fingerprint(article: Article) -> Fingerprint: title=article.title, publish_date=_publish_date(article), ingested_at=datetime.now(), + source_ids=[article.source_id], ) @@ -80,7 +84,7 @@ class Deduper: # ------------------------------------------------------------------ # def check(self, article: Article) -> DedupResult: - """三层判重(只读)。""" + """三层判重(只读)。命中时附带匹配指纹的多源信息(all_source_ids)。""" fp = article_to_fingerprint(article) # L1: URL hash @@ -93,6 +97,8 @@ class Deduper: matched_url_hash=existing.url_hash, matched_url=existing.url, matched_title=existing.title, + matched_source_id=existing.source_id, + all_source_ids=existing.source_ids, ) # L2: 内容 hash @@ -105,6 +111,8 @@ class Deduper: matched_url_hash=existing.url_hash, matched_url=existing.url, matched_title=existing.title, + matched_source_id=existing.source_id, + all_source_ids=existing.source_ids, ) # L3: SimHash 模糊 @@ -129,20 +137,40 @@ class Deduper: matched_url_hash=best_match.url_hash, matched_url=best_match.url, matched_title=best_match.title, + matched_source_id=best_match.source_id, + all_source_ids=best_match.source_ids, hamming_distance=best_dist, ) return DedupResult(url_hash=fp.url_hash, is_duplicate=False) def ingest(self, article: Article) -> DedupResult: - """判重 + 不重复则入库。""" + """判重 + 不重复则入库。 + + 命中重复时,把当前文章的 source_id 合并进匹配指纹的 source_ids + (记录同一内容组的全部来源),并更新 all_source_ids 后返回。 + """ result = self.check(article) if not result.is_duplicate: fp = article_to_fingerprint(article) self.store.upsert(fp) logger.debug("入库: {} {}", fp.url_hash, fp.title[:30]) - else: - logger.debug("命中重复: {}", result.short_summary()) + return result + + # 重复:合并来源到匹配指纹(主源保持首位,Fingerprint validator 负责去重) + if result.matched_url_hash and article.source_id not in result.all_source_ids: + matched = self.store.get_by_url_hash(result.matched_url_hash) + if matched is not None: + merged = [*matched.source_ids, article.source_id] + self.store.upsert(matched.model_copy(update={"source_ids": merged})) + result = result.model_copy( + update={"all_source_ids": merged} + ) + logger.debug( + "合并来源 {} -> {} ({} 个源)", + article.source_id, result.matched_url_hash, len(merged), + ) + logger.debug("命中重复: {}", result.short_summary()) return result def stats(self) -> DedupStats: diff --git a/dedup/models.py b/dedup/models.py index 9ad6b22..99dc1b7 100644 --- a/dedup/models.py +++ b/dedup/models.py @@ -4,9 +4,9 @@ from __future__ import annotations from datetime import datetime from enum import StrEnum -from typing import Literal +from typing import Literal, Self -from pydantic import BaseModel, Field +from pydantic import BaseModel, Field, model_validator class DedupLayer(StrEnum): @@ -18,7 +18,12 @@ class DedupLayer(StrEnum): class Fingerprint(BaseModel): - """单篇文章的指纹记录,持久化到 SQLite。""" + """单篇文章的指纹记录,持久化到 SQLite。 + + source_ids: 同一内容组(去重后视为同一篇新闻)的全部来源列表, + 第一位是主源(即本指纹的 source_id);重复文章命中时由 + Deduper.ingest 自动合并,实现「一条唯一新闻记录多个源」。 + """ url_hash: str = Field(..., description="主键,与 Article.url_hash 一致") content_hash: str = Field(..., description="标准化 content 的 SHA1[:16]") @@ -28,6 +33,20 @@ class Fingerprint(BaseModel): title: str publish_date: str | None = Field(default=None, description="YYYY-MM-DD,用于时间窗口") ingested_at: datetime = Field(default_factory=datetime.now) + source_ids: list[str] = Field( + default_factory=list, + description="同内容组全部来源(去重合并),始终包含 source_id 且其居首", + ) + + @model_validator(mode="after") + def _ensure_source_ids(self) -> Self: + """保证 source_ids 非空、去重且以主源 source_id 开头。""" + seen: list[str] = [] + for s in [self.source_id, *self.source_ids]: + if s and s not in seen: + seen.append(s) + self.source_ids = seen + return self class DedupResult(BaseModel): @@ -39,6 +58,13 @@ class DedupResult(BaseModel): matched_url_hash: str | None = None matched_url: str | None = None matched_title: str | None = None + matched_source_id: str | None = Field( + default=None, description="匹配指纹的主源 source_id" + ) + all_source_ids: list[str] = Field( + default_factory=list, + description="该内容组(唯一新闻)的全部来源;含匹配指纹自身的来源", + ) hamming_distance: int | None = Field( default=None, description="仅 SimHash 层有值" ) diff --git a/dedup/store.py b/dedup/store.py index 9e8b52c..f6edd51 100644 --- a/dedup/store.py +++ b/dedup/store.py @@ -3,10 +3,14 @@ 注意:SimHash 是 64 位无符号整数,SQLite INTEGER 是 64 位有符号 (范围 [-2^63, 2^63-1])。直接存可能溢出/转负数,虽然 XOR 仍然 正确但语义混乱。这里统一存为 16 位 hex TEXT,避免符号问题。 + +source_ids 列存 JSON 数组文本(同一内容组全部来源);旧库无此列时 +自动 ALTER TABLE 迁移,旧数据读取时回退为 [source_id]。 """ from __future__ import annotations +import json import sqlite3 from datetime import datetime, timedelta from pathlib import Path @@ -27,13 +31,17 @@ CREATE TABLE IF NOT EXISTS fingerprints ( url TEXT NOT NULL, title TEXT NOT NULL, publish_date TEXT, - ingested_at TEXT NOT NULL + ingested_at TEXT NOT NULL, + source_ids TEXT ); CREATE INDEX IF NOT EXISTS idx_content_hash ON fingerprints(content_hash); CREATE INDEX IF NOT EXISTS idx_publish_date ON fingerprints(publish_date); CREATE INDEX IF NOT EXISTS idx_source_id ON fingerprints(source_id); """ +# 兼容旧库:为已存在但缺少 source_ids 列的表补列 +_ALTER_SQL = "ALTER TABLE fingerprints ADD COLUMN source_ids TEXT" + def _to_hex(simhash: int) -> str: return f"{simhash:016x}" @@ -43,6 +51,25 @@ def _from_hex(hex_str: str) -> int: return int(hex_str, 16) +def _to_sources_json(source_ids: list[str]) -> str: + return json.dumps(source_ids, ensure_ascii=False) + + +def _from_sources_json(raw: str | None, fallback: str) -> list[str]: + """解析 source_ids 列;NULL/损坏时回退 [主源]。""" + if not raw: + return [fallback] + try: + val = json.loads(raw) + except (TypeError, ValueError): + return [fallback] + if isinstance(val, list) and val: + # 保证主源在首位(兼容手改/旧数据) + cleaned = [s for s in val if s and s != fallback] + return [fallback, *cleaned] + return [fallback] + + def _row_to_fp(row: sqlite3.Row) -> Fingerprint: return Fingerprint( url_hash=row["url_hash"], @@ -53,6 +80,7 @@ def _row_to_fp(row: sqlite3.Row) -> Fingerprint: title=row["title"], publish_date=row["publish_date"], ingested_at=datetime.fromisoformat(row["ingested_at"]), + source_ids=_from_sources_json(row["source_ids"], row["source_id"]), ) @@ -67,8 +95,17 @@ class FingerprintStore: ) self._conn.row_factory = sqlite3.Row self._conn.executescript(_SCHEMA_SQL) + self._migrate_source_ids() logger.debug("打开指纹库: {}", self.db_path) + def _migrate_source_ids(self) -> None: + """旧库兼容:为缺少 source_ids 列的表补列(新库无需执行)。""" + try: + self._conn.execute(_ALTER_SQL) + logger.info("指纹库迁移:为 fingerprints 表新增 source_ids 列") + except sqlite3.OperationalError: + logger.debug("source_ids 列已存在,跳过迁移") + def close(self) -> None: self._conn.close() @@ -148,7 +185,7 @@ class FingerprintStore: self._conn.execute( "INSERT OR REPLACE INTO fingerprints " "(url_hash, content_hash, simhash_hex, source_id, url, title, " - " publish_date, ingested_at) VALUES (?, ?, ?, ?, ?, ?, ?, ?)", + " publish_date, ingested_at, source_ids) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?)", ( fp.url_hash, fp.content_hash, @@ -158,6 +195,7 @@ class FingerprintStore: fp.title, fp.publish_date, fp.ingested_at.isoformat(), + _to_sources_json(fp.source_ids), ), ) diff --git a/embedding/factory.py b/embedding/factory.py index 0d7a815..10f7c8f 100644 --- a/embedding/factory.py +++ b/embedding/factory.py @@ -1,9 +1,14 @@ -"""Embedding provider 工厂:根据环境变量构造合适后端。""" +"""Embedding provider 工厂:根据配置构造合适后端。 + +配置优先级: 显式参数 > configs/llm_models.yaml scenes.embedding > .env > 默认。 +""" from __future__ import annotations import os +from configs.loader import load_scene_config + from .base import AsyncEmbeddingProvider, EmbeddingProvider from .models import EmbeddingError, EmbeddingProviderType from .remote import ( @@ -20,14 +25,20 @@ def _read_env(key: str, default: str | None = None) -> str | None: def resolve_provider_type(provider: str | None = None) -> EmbeddingProviderType: - """根据 provider 参数 / env 解析出 EmbeddingProviderType。 + """根据 provider 参数 / YAML 场景 / env 解析出 EmbeddingProviderType。 映射: dashscope / qwen / remote -> DASHSCOPE local / local-bge / bge / bge-m3 -> LOCAL_BGE 默认 dashscope。 """ - p = (provider or _read_env("EMBEDDING_PROVIDER", "dashscope") or "dashscope").lower() + scene_provider = load_scene_config("embedding").get("provider") + p = ( + provider + or scene_provider + or _read_env("EMBEDDING_PROVIDER", "dashscope") + or "dashscope" + ).lower() if p in ("dashscope", "qwen", "remote"): return EmbeddingProviderType.DASHSCOPE if p in ("local", "local-bge", "bge", "bge-m3"): diff --git a/embedding/local.py b/embedding/local.py index feaad4f..4aff44f 100644 --- a/embedding/local.py +++ b/embedding/local.py @@ -21,6 +21,8 @@ from .models import EmbeddingError if TYPE_CHECKING: from sentence_transformers import SentenceTransformer +from configs.loader import load_scene_config + LOCAL_DEFAULT_MODEL = "BAAI/bge-m3" LOCAL_DEFAULT_DIM = 1024 @@ -56,10 +58,12 @@ class LocalBGEEmbeddingProvider(EmbeddingProvider): normalize: bool = True, ) -> None: st_cls = _try_import_st() - # 模型优先级:CLI/参数 > LOCAL_EMBEDDING_MODEL > 默认值 + # 模型优先级:CLI/参数 > YAML scenes.embedding > LOCAL_EMBEDDING_MODEL > 默认值 # 不读全局 EMBEDDING_MODEL,避免与 DashScope 冲突 + scene_model = load_scene_config("embedding").get("model") self.model = ( model + or scene_model or _read_env("LOCAL_EMBEDDING_MODEL", LOCAL_DEFAULT_MODEL) or LOCAL_DEFAULT_MODEL ) diff --git a/embedding/remote.py b/embedding/remote.py index da74bc4..36d1b6e 100644 --- a/embedding/remote.py +++ b/embedding/remote.py @@ -5,10 +5,14 @@ model: text-embedding-v3 (1024 维) 限制: 单次请求 input ≤ 25 条 -环境变量: - DASHSCOPE_API_KEY - QWEN_BASE_URL (默认百炼兼容路径) - EMBEDDING_MODEL (默认 text-embedding-v3) +配置来源(优先级从高到低): + 1. 构造参数(model / api_key / base_url / max_attempts) + 2. configs/llm_models.yaml 的 scenes.embedding + 3. 环境变量 / .env: + DASHSCOPE_EMBEDDING_API_KEY / DASHSCOPE_API_KEY + DASHSCOPE_EMBEDDING_BASE_URL / QWEN_BASE_URL + DASHSCOPE_EMBEDDING_MODEL + 4. 代码内置默认值(text-embedding-v3 / 1024 维) """ from __future__ import annotations @@ -19,6 +23,8 @@ import os from loguru import logger from openai import AsyncOpenAI, OpenAI +from configs.loader import load_scene_config + from .base import AsyncEmbeddingProvider, EmbeddingProvider from .models import EmbeddingError @@ -32,6 +38,9 @@ DEFAULT_MAX_ATTEMPTS = 3 RETRY_BASE_WAIT_SEC = 1.0 RETRY_MAX_WAIT_SEC = 8.0 +# embedding 场景名(对应 configs/llm_models.yaml scenes.embedding) +SCENE_EMBEDDING = "embedding" + def _read_env(key: str, default: str | None = None) -> str | None: val = os.environ.get(key) @@ -40,23 +49,48 @@ def _read_env(key: str, default: str | None = None) -> str | None: return val.strip() +def _scene() -> dict: + """读取 YAML embedding 场景配置(不存在时为空 dict)。""" + return load_scene_config(SCENE_EMBEDDING) + + +def _scene_int(key: str, default: int) -> int: + try: + return int(_scene().get(key) or default) + except (TypeError, ValueError): + return default + + +def _scene_float(key: str, default: float) -> float: + try: + return float(_scene().get(key) or default) + except (TypeError, ValueError): + return default + + def _resolve_config() -> tuple[str, str, str]: """读取 API key / base_url / model,返回 (api_key, base_url, model)。 + 优先级: YAML 场景 > 环境变量 > 内置默认。 模型名优先级:DASHSCOPE_EMBEDDING_MODEL > 默认值。 不再读全局 EMBEDDING_MODEL,避免与 LOCAL provider 冲突。 """ - api_key = _read_env("DASHSCOPE_EMBEDDING_API_KEY") or _read_env("DASHSCOPE_API_KEY") or "" + sc = _scene() + api_key_env = sc.get("api_key_env") or "DASHSCOPE_EMBEDDING_API_KEY" + api_key = _read_env(api_key_env) or _read_env("DASHSCOPE_API_KEY") or "" if not api_key: - raise EmbeddingError("DASHSCOPE_EMBEDDING_API_KEY 或 DASHSCOPE_API_KEY 未配置") - # DASHSCOPE_EMBEDDING_BASE_URL -> QWEN_BASE_URL(兜底) -> 默认 + raise EmbeddingError( + f"{api_key_env} 或 DASHSCOPE_API_KEY 未配置" + ) + # YAML base_url_env -> DASHSCOPE_EMBEDDING_BASE_URL -> QWEN_BASE_URL(兜底) -> 默认 base_url = ( - _read_env("DASHSCOPE_EMBEDDING_BASE_URL") + _read_env(sc.get("base_url_env") or "DASHSCOPE_EMBEDDING_BASE_URL") or _read_env("QWEN_BASE_URL") or DASHSCOPE_DEFAULT_BASE ) model = ( - _read_env("DASHSCOPE_EMBEDDING_MODEL", DASHSCOPE_DEFAULT_MODEL) + sc.get("model") + or _read_env("DASHSCOPE_EMBEDDING_MODEL", DASHSCOPE_DEFAULT_MODEL) or DASHSCOPE_DEFAULT_MODEL ) return api_key, base_url, model @@ -78,24 +112,27 @@ class DashScopeEmbeddingProvider(EmbeddingProvider): model: str | None = None, api_key: str | None = None, base_url: str | None = None, - timeout_sec: float = 60.0, - max_attempts: int = DEFAULT_MAX_ATTEMPTS, + timeout_sec: float | None = None, + max_attempts: int | None = None, + batch_limit: int | None = None, ) -> None: env_key, env_base, env_model = _resolve_config() self.model = model or env_model self.dim = DASHSCOPE_DEFAULT_DIM - self.max_attempts = max_attempts + self.max_attempts = max_attempts or _scene_int("max_attempts", DEFAULT_MAX_ATTEMPTS) + self.batch_limit = batch_limit or _scene_int("batch_limit", DASHSCOPE_BATCH_LIMIT) + timeout = timeout_sec or _scene_float("timeout_sec", 60.0) self._client = OpenAI( api_key=api_key or env_key, base_url=base_url or env_base, - timeout=timeout_sec, + timeout=timeout, ) def embed_batch(self, texts: list[str]) -> list[list[float]]: if not texts: return [] results: list[list[float]] = [] - for chunk in _chunked(texts, DASHSCOPE_BATCH_LIMIT): + for chunk in _chunked(texts, self.batch_limit): results.extend(self._call_with_retry(chunk)) return results @@ -136,24 +173,27 @@ class DashScopeAsyncEmbeddingProvider(AsyncEmbeddingProvider): model: str | None = None, api_key: str | None = None, base_url: str | None = None, - timeout_sec: float = 60.0, - max_attempts: int = DEFAULT_MAX_ATTEMPTS, + timeout_sec: float | None = None, + max_attempts: int | None = None, + batch_limit: int | None = None, ) -> None: env_key, env_base, env_model = _resolve_config() self.model = model or env_model self.dim = DASHSCOPE_DEFAULT_DIM - self.max_attempts = max_attempts + self.max_attempts = max_attempts or _scene_int("max_attempts", DEFAULT_MAX_ATTEMPTS) + self.batch_limit = batch_limit or _scene_int("batch_limit", DASHSCOPE_BATCH_LIMIT) + timeout = timeout_sec or _scene_float("timeout_sec", 60.0) self._client = AsyncOpenAI( api_key=api_key or env_key, base_url=base_url or env_base, - timeout=timeout_sec, + timeout=timeout, ) async def embed_batch(self, texts: list[str]) -> list[list[float]]: if not texts: return [] results: list[list[float]] = [] - for chunk in _chunked(texts, DASHSCOPE_BATCH_LIMIT): + for chunk in _chunked(texts, self.batch_limit): results.extend(await self._call_with_retry(chunk)) return results diff --git a/llm/__init__.py b/llm/__init__.py index a05a865..8ac4347 100644 --- a/llm/__init__.py +++ b/llm/__init__.py @@ -8,15 +8,18 @@ """ from .client import ( + DEFAULT_MAX_ATTEMPTS, DEFAULT_TEMPERATURE, DEFAULT_TIMEOUT_SEC, + SCENE_DAILY_REPORT, + SCENE_EVENT_EXTRACTION, + SCENE_STOCK_REPORT, LLMConfig, load_llm_config, make_async_client, make_sync_client, ) from .extractor import ( - DEFAULT_MAX_ATTEMPTS, DEFAULT_PROMPT_PATH, MAX_CONTENT_CHARS, PromptTemplate, @@ -43,6 +46,9 @@ __all__ = [ "MAX_CONTENT_CHARS", "MAX_IMPORTANCE", "MIN_IMPORTANCE", + "SCENE_DAILY_REPORT", + "SCENE_EVENT_EXTRACTION", + "SCENE_STOCK_REPORT", "EventExtraction", "ExtractedEvent", "LLMCallError", diff --git a/llm/client.py b/llm/client.py index 3396721..7f60ec5 100644 --- a/llm/client.py +++ b/llm/client.py @@ -2,7 +2,13 @@ 支持 DeepSeek 和 Qwen(百炼),两者均为 OpenAI 兼容接口,共用 openai SDK。 -环境变量: +配置来源(优先级从高到低): + 1. 代码 / CLI 显式参数(provider / model) + 2. configs/llm_models.yaml 场景配置(scene 参数,见 configs/loader.py) + 3. 环境变量 / .env(LLM_PROVIDER、DEEPSEEK_MODEL 等,向后兼容) + 4. 代码内置默认值 + +环境变量(兜底): LLM_PROVIDER = deepseek | qwen (默认 deepseek) DeepSeek: DEEPSEEK_API_KEY / DEEPSEEK_BASE_URL / DEEPSEEK_MODEL Qwen: QWEN_API_KEY / QWEN_BASE_URL / QWEN_MODEL @@ -19,6 +25,8 @@ from dataclasses import dataclass from loguru import logger from openai import AsyncOpenAI, OpenAI +from configs.loader import load_defaults, load_scene_config + # 默认基址 _DEEPSEEK_DEFAULT_BASE = "https://api.deepseek.com" _QWEN_DEFAULT_BASE = "https://dashscope.aliyuncs.com/compatible-mode/v1" @@ -26,6 +34,12 @@ _QWEN_DEFAULT_BASE = "https://dashscope.aliyuncs.com/compatible-mode/v1" # 抽取任务默认参数 DEFAULT_TIMEOUT_SEC = 60.0 DEFAULT_TEMPERATURE = 0.1 +DEFAULT_MAX_ATTEMPTS = 3 + +# 场景名 -> configs/llm_models.yaml 中 scenes 的 key +SCENE_EVENT_EXTRACTION = "event_extraction" +SCENE_DAILY_REPORT = "daily_report" +SCENE_STOCK_REPORT = "stock_report" @dataclass @@ -38,6 +52,7 @@ class LLMConfig: base_url: str timeout_sec: float = DEFAULT_TIMEOUT_SEC temperature: float = DEFAULT_TEMPERATURE + max_attempts: int = DEFAULT_MAX_ATTEMPTS # 单次任务失败重试次数 def __post_init__(self) -> None: if not self.api_key: @@ -51,49 +66,122 @@ def _read_env(key: str, default: str | None = None) -> str | None: return val.strip() +def _first_env(keys: list[str | None]) -> str | None: + """按顺序返回第一个非空的环境变量值。""" + for k in keys: + if not k: + continue + v = _read_env(k) + if v: + return v + return None + + +def _num(value: object) -> float | None: + """把 YAML 数字/字符串安全转 float;非法或为空返回 None。""" + if value is None or value == "": + return None + try: + return float(value) + except (TypeError, ValueError): + return None + + def load_llm_config( provider: str | None = None, *, model: str | None = None, + scene: str | None = None, ) -> LLMConfig: - """根据环境变量构造 LLMConfig。 + """按优先级构造 LLMConfig:显式参数 > YAML 场景 > 环境变量 > 内置默认。 - provider 为 None 时读 LLM_PROVIDER 环境变量,默认 deepseek。 - model 为 None 时读 LLM_MODEL 或 provider 默认。 + scene 对应 configs/llm_models.yaml 中 scenes 的 key + (event_extraction / daily_report / stock_report),该场景未配置的字段 + 回退到环境变量,保持向后兼容。 """ - p = (provider or _read_env("LLM_PROVIDER", "deepseek") or "deepseek").lower() + sc = load_scene_config(scene or "") + dflt = load_defaults() + + p = ( + provider + or sc.get("provider") + or _read_env("LLM_PROVIDER", "deepseek") + or "deepseek" + ).lower() + + # 各 provider 的 api_key / base_url / model 环境变量链 + provider_envs: dict[str, tuple[list[str | None], list[str | None], list[str | None]]] = { + "deepseek": ( + [sc.get("api_key_env"), "DEEPSEEK_API_KEY"], + [sc.get("base_url_env"), "DEEPSEEK_BASE_URL"], + ["DEEPSEEK_MODEL", "LLM_MODEL"], + ), + "qwen": ( + [sc.get("api_key_env"), "QWEN_API_KEY", "DASHSCOPE_API_KEY"], + [sc.get("base_url_env"), "QWEN_BASE_URL"], + ["QWEN_MODEL", "LLM_MODEL"], + ), + } if p == "deepseek": - api_key = _read_env("DEEPSEEK_API_KEY") or "" - base = _read_env("DEEPSEEK_BASE_URL", _DEEPSEEK_DEFAULT_BASE) or _DEEPSEEK_DEFAULT_BASE - # DEEPSEEK_MODEL → LLM_MODEL;模型必须显式配置,不提供内置默认 - m = model or _read_env("DEEPSEEK_MODEL") or _read_env("LLM_MODEL") - if not m: - raise ValueError("未配置 LLM 模型: 请设置 DEEPSEEK_MODEL 或 LLM_MODEL") + key_envs, base_envs, model_envs = provider_envs["deepseek"] + default_base = _DEEPSEEK_DEFAULT_BASE elif p in ("qwen", "dashscope"): - api_key = _read_env("QWEN_API_KEY") or _read_env("DASHSCOPE_API_KEY") or "" - base = _read_env("QWEN_BASE_URL", _QWEN_DEFAULT_BASE) or _QWEN_DEFAULT_BASE - # QWEN_MODEL → LLM_MODEL;模型必须显式配置,不提供内置默认 - m = model or _read_env("QWEN_MODEL") or _read_env("LLM_MODEL") - if not m: - raise ValueError("未配置 LLM 模型: 请设置 QWEN_MODEL 或 LLM_MODEL") + key_envs, base_envs, model_envs = provider_envs["qwen"] + default_base = _QWEN_DEFAULT_BASE p = "qwen" # 内部统一用 qwen else: raise ValueError(f"未知 LLM provider: {p!r},仅支持 deepseek / qwen") - timeout = float(_read_env("LLM_TIMEOUT_SEC", str(DEFAULT_TIMEOUT_SEC)) or DEFAULT_TIMEOUT_SEC) - temperature = float(_read_env("LLM_TEMPERATURE", str(DEFAULT_TEMPERATURE)) or DEFAULT_TEMPERATURE) + api_key = _first_env(key_envs) or "" + base_url = _first_env(base_envs) or default_base + # 模型优先级:显式参数 > YAML 场景 > 环境变量;模型必须显式配置,无内置兜底 + m = model or sc.get("model") or _first_env(model_envs) + if not m: + env_hint = "/".join(v for v in model_envs if v) + raise ValueError( + f"未配置 LLM 模型(场景 {scene or 'default'}): " + f"请在 configs/llm_models.yaml 的 model 或 .env 设置 {env_hint}" + ) + + timeout = _pick_float(sc, dflt, "timeout_sec", "LLM_TIMEOUT_SEC", DEFAULT_TIMEOUT_SEC) + temperature = _pick_float(sc, dflt, "temperature", "LLM_TEMPERATURE", DEFAULT_TEMPERATURE) + max_attempts = _pick_int(sc, "max_attempts", DEFAULT_MAX_ATTEMPTS) return LLMConfig( provider=p, model=m, api_key=api_key, - base_url=base, + base_url=base_url, timeout_sec=timeout, temperature=temperature, + max_attempts=max_attempts, ) +def _pick_float( + sc: dict, + dflt: dict, + sc_key: str, + env_key: str, + default: float, +) -> float: + """数值参数选择:YAML 场景 > 环境变量 > YAML defaults > 内置默认(零值合法)。""" + v = _num(sc.get(sc_key)) + if v is not None: + return v + v = _num(_read_env(env_key)) + if v is not None: + return v + v = _num(dflt.get(sc_key)) + return v if v is not None else default + + +def _pick_int(sc: dict, sc_key: str, default: int) -> int: + v = _num(sc.get(sc_key)) + return int(v) if v is not None else default + + def make_sync_client(config: LLMConfig) -> OpenAI: """构造同步 OpenAI 客户端(指向 DeepSeek/Qwen 兼容端点)。""" logger.debug( diff --git a/llm/extractor.py b/llm/extractor.py index 8dea281..3b09fba 100644 --- a/llm/extractor.py +++ b/llm/extractor.py @@ -187,11 +187,15 @@ def extract_event( article: Article, *, template: PromptTemplate | None = None, - max_attempts: int = DEFAULT_MAX_ATTEMPTS, + max_attempts: int | None = None, ) -> ExtractedEvent: - """同步抽取单篇文章的事件(带重试)。""" + """同步抽取单篇文章的事件(带重试)。 + + max_attempts 为 None 时使用 config.max_attempts(来自 YAML/环境变量配置)。 + """ tpl = template or PromptTemplate() prompt = tpl.render(article) + max_attempts = max_attempts or config.max_attempts last_err: Exception | None = None for attempt in range(1, max_attempts + 1): @@ -241,12 +245,16 @@ async def extract_event_async( article: Article, *, template: PromptTemplate | None = None, - max_attempts: int = DEFAULT_MAX_ATTEMPTS, + max_attempts: int | None = None, semaphore: asyncio.Semaphore | None = None, ) -> ExtractedEvent: - """异步抽取(批处理用),与同步版逻辑等价。""" + """异步抽取(批处理用),与同步版逻辑等价。 + + max_attempts 为 None 时使用 config.max_attempts。 + """ tpl = template or PromptTemplate() prompt = tpl.render(article) + max_attempts = max_attempts or config.max_attempts async def _run() -> ExtractedEvent: last_err: Exception | None = None diff --git a/scheduler/reporter.py b/scheduler/reporter.py index 7c1583e..607721c 100644 --- a/scheduler/reporter.py +++ b/scheduler/reporter.py @@ -17,7 +17,10 @@ import time from collections import Counter from datetime import date, datetime, timedelta from pathlib import Path -from typing import Any +from typing import TYPE_CHECKING, Any + +if TYPE_CHECKING: + from llm.client import LLMConfig from dotenv import load_dotenv from loguru import logger @@ -519,10 +522,10 @@ def _generate_ai_summary(news: dict, cninfo: dict, day_str: str, return "" try: - from llm.client import load_llm_config, make_sync_client - config = load_llm_config() + from llm.client import SCENE_DAILY_REPORT, load_llm_config, make_sync_client + config = load_llm_config(scene=SCENE_DAILY_REPORT) client = make_sync_client(config) - return _llm_summarize(client, config.model, lines, day_str) + return _llm_summarize(client, config, lines, day_str) except Exception as e: logger.warning("AI 摘要生成失败: {}", e) return "" @@ -548,8 +551,12 @@ def _split_lines_into_chunks(lines: list[str], max_chars: int = 3000) -> list[li return chunks -def _llm_summarize(client, model: str, lines: list[str], day_str: str) -> str: - """LLM 摘要:单块直接总结,多块先分段总结再合并。""" +def _llm_summarize(client, config: LLMConfig, lines: list[str], day_str: str) -> str: + """LLM 摘要:单块直接总结,多块先分段总结再合并。 + + config 为 llm.client.LLMConfig(daily_report 场景),提供 model / temperature。 + """ + model = config.model chunks = _split_lines_into_chunks(lines) if len(chunks) == 1: @@ -609,9 +616,10 @@ def _build_prompt(lines: list[str], day_str: str) -> str: 直接输出要点列表:""" -def _llm_call(client, model: str, prompt: str, max_tokens: int = 1500) -> str: +def _llm_call(client, config: LLMConfig, prompt: str, max_tokens: int = 1500) -> str: """单次 LLM 调用(带重试),返回 strip 后的文本。 + config 为 llm.client.LLMConfig(daily_report 场景),提供 model / temperature。 失败按指数退避重试 `_LLM_RETRY_TIMES` 次(默认 3),全部失败则抛出最后一次异常。 若 finish_reason 为 'length' 则说明达到 max_tokens 上限被截断。 """ @@ -619,12 +627,12 @@ def _llm_call(client, model: str, prompt: str, max_tokens: int = 1500) -> str: for attempt in range(_LLM_RETRY_TIMES): try: resp = client.chat.completions.create( - model=model, + model=config.model, messages=[ {"role": "system", "content": "你是 A 股日报撰写助手,输出简洁、有洞察的新闻摘要。"}, {"role": "user", "content": prompt}, ], - temperature=0.3, + temperature=config.temperature, max_tokens=max_tokens, ) content = (resp.choices[0].message.content or "").strip() diff --git a/scheduler/stock_reporter.py b/scheduler/stock_reporter.py index 510d596..b6e57b4 100644 --- a/scheduler/stock_reporter.py +++ b/scheduler/stock_reporter.py @@ -232,7 +232,7 @@ def _generate_ai_summary(company_name: str, announcements: list[dict], news: list[dict], research: list[dict], irm: list[dict]) -> str: """LLM 生成个股要点分析。""" - from llm.client import load_llm_config, make_sync_client + from llm.client import SCENE_STOCK_REPORT, load_llm_config, make_sync_client lines = [] @@ -277,12 +277,12 @@ def _generate_ai_summary(company_name: str, announcements: list[dict], 直接输出要点列表:""" try: - config = load_llm_config() + config = load_llm_config(scene=SCENE_STOCK_REPORT) client = make_sync_client(config) resp = client.chat.completions.create( model=config.model, messages=[{"role": "user", "content": prompt}], - temperature=0.3, max_tokens=500, + temperature=config.temperature, max_tokens=500, ) return (resp.choices[0].message.content or "").strip() except Exception as e: diff --git a/scripts/run_dedup.py b/scripts/run_dedup.py index a0ace13..c84e93d 100644 --- a/scripts/run_dedup.py +++ b/scripts/run_dedup.py @@ -2,9 +2,12 @@ 输入: data/processed/{source}/{YYYYMMDD}/*.json (M2 产物) 输出: - - 指纹库:data/dedup/fingerprints.sqlite3 - - 唯一文章:data/deduped/{YYYYMMDD}/uniques/{url_hash}.json + - 指纹库:data/dedup/fingerprints.sqlite3 (source_ids 列记录多源) + - 唯一文章:data/deduped/{YYYYMMDD}/uniques/{url_hash}.json (含 sources 多源字段) + - 多源记录:data/deduped/{YYYYMMDD}/sources.json + {url_hash: [source_id, ...]},一条唯一新闻的全部来源 - 重复记录:data/deduped/{YYYYMMDD}/duplicates.jsonl + (含 matched_source_id / matched_source_ids) 用法: uv run python -m scripts.run_dedup # 处理今日全部源 @@ -56,14 +59,60 @@ def _list_source_dirs(processed_root: Path) -> list[str]: return sorted(p.name for p in processed_root.iterdir() if p.is_dir()) +def _write_unique( + url_hash: str, + article: Article, + uniques_dir: Path, + sources_map: dict[str, list[str]], +) -> None: + """写 uniques JSON,附加 sources 多源字段(向后兼容:下游 Pydantic 忽略多余字段)。""" + data = json.loads(article.model_dump_json()) + data["sources"] = list(dict.fromkeys(sources_map.get(url_hash, [article.source_id]))) + out_path = uniques_dir / f"{url_hash}.json" + out_path.write_text(json.dumps(data, ensure_ascii=False, indent=2), encoding="utf-8") + + +def _update_unique_sources( + url_hash: str, + uniques_dir: Path, + sources_map: dict[str, list[str]], +) -> None: + """仅更新已存在 uniques 文件的 sources 字段(不覆盖原文内容)。 + + 跨日命中时对应 uniques 文件在历史日期目录,不在本次处理范围,以指纹库为准。 + """ + uniq_path = uniques_dir / f"{url_hash}.json" + if not uniq_path.is_file(): + return + try: + data = json.loads(uniq_path.read_text(encoding="utf-8")) + data["sources"] = list(dict.fromkeys(sources_map.get(url_hash, []))) + uniq_path.write_text(json.dumps(data, ensure_ascii=False, indent=2), encoding="utf-8") + except (json.JSONDecodeError, OSError) as e: + logger.warning("更新 uniques 多源失败 {}: {}", uniq_path, e) + + +def _merge_sources(url_hash: str, new_source: str, sources_map: dict[str, list[str]]) -> None: + """把新源并入 url_hash 的源列表(去重保序,主源居首)。""" + cur = sources_map.setdefault(url_hash, []) + if new_source not in cur: + cur.append(new_source) + + def _process_source_day( source_id: str, day: str, processed_root: Path, out_root: Path, deduper: Deduper, + sources_map: dict[str, list[str]], ) -> tuple[int, int, Counter]: - """处理单源单日。返回 (uniques, duplicates, layer_counter)。""" + """处理单源单日。返回 (uniques, duplicates, layer_counter)。 + + sources_map: 本次去重涉及内容组的 url_hash -> 全部来源列表(跨源累积, + 最终写入 data/deduped/{day}/sources.json,供「显示新闻源」使用; + 跨日命中的历史内容组也会记录,权威多源以指纹库 source_ids 列为准)。 + """ src_dir = processed_root / source_id / day if not src_dir.is_dir(): logger.info("源 {} 日期 {} 无 processed 目录,跳过", source_id, day) @@ -92,6 +141,11 @@ def _process_source_day( dup_cnt += 1 if result.matched_layer is not None: layer_cnt[result.matched_layer.value] += 1 + # 记录多源:把被去重文章的源并入对应唯一新闻 + if result.matched_url_hash: + _merge_sources(result.matched_url_hash, article.source_id, sources_map) + # 若该唯一新闻文件在当天目录,同步更新其 sources 字段 + _update_unique_sources(result.matched_url_hash, uniques_dir, sources_map) dup_f.write( json.dumps( { @@ -107,6 +161,8 @@ def _process_source_day( "matched_url": result.matched_url, "matched_url_hash": result.matched_url_hash, "matched_title": result.matched_title, + "matched_source_id": result.matched_source_id, + "matched_source_ids": result.all_source_ids, "hamming_distance": result.hamming_distance, }, ensure_ascii=False, @@ -115,8 +171,8 @@ def _process_source_day( ) else: uniq_cnt += 1 - out_path = uniques_dir / f"{article.url_hash}.json" - out_path.write_text(article.model_dump_json(indent=2), encoding="utf-8") + sources_map[article.url_hash] = [article.source_id] + _write_unique(article.url_hash, article, uniques_dir, sources_map) total = uniq_cnt + dup_cnt rate = dup_cnt / max(total, 1) @@ -173,14 +229,35 @@ def main() -> int: total_uniq = 0 total_dup = 0 total_layers: Counter = Counter() + # 当天唯一新闻 url_hash -> 全部来源列表(跨源累积,多源记录) + sources_map: dict[str, list[str]] = {} for src in sources: u, d, lc = _process_source_day( - src, args.date, processed_root, out_root, deduper + src, args.date, processed_root, out_root, deduper, sources_map ) total_uniq += u total_dup += d total_layers.update(lc) + # 多源记录汇总:data/deduped/{day}/sources.json + # {url_hash: [source_id, ...]},配合 uniques/{url_hash}.json 的 sources 字段 + # 与指纹库 source_ids 列,提供「一条唯一新闻多个来源」的完整记录。 + sources_path = out_root / args.date / "sources.json" + sources_path.write_text( + json.dumps( + {k: v for k, v in sources_map.items() if v}, + ensure_ascii=False, + indent=2, + ), + encoding="utf-8", + ) + logger.info( + "多源记录已写入 {} ({} 条唯一新闻,{} 条含多源)", + sources_path, + len(sources_map), + sum(1 for v in sources_map.values() if len(v) > 1), + ) + total = total_uniq + total_dup rate = total_dup / max(total, 1) logger.info( diff --git a/scripts/run_event_extraction.py b/scripts/run_event_extraction.py index ac91748..c7de976 100644 --- a/scripts/run_event_extraction.py +++ b/scripts/run_event_extraction.py @@ -31,6 +31,7 @@ from pydantic import ValidationError from extractor import Article from llm import ( + SCENE_EVENT_EXTRACTION, ExtractedEvent, LLMCallError, PromptTemplate, @@ -88,7 +89,10 @@ def _load_article(p: Path) -> Article | None: async def _run(args: argparse.Namespace) -> int: load_dotenv() # 读 .env 到 os.environ - config = load_llm_config(provider=args.provider, model=args.model) + # scene=event_extraction: 读取 configs/llm_models.yaml 场景 1 配置,未配置字段回退 .env + config = load_llm_config( + provider=args.provider, model=args.model, scene=SCENE_EVENT_EXTRACTION + ) logger.info( "LLM provider={} model={} base_url={}", config.provider, config.model, config.base_url, diff --git a/tests/test_dedup.py b/tests/test_dedup.py index 0a2016e..69c07e6 100644 --- a/tests/test_dedup.py +++ b/tests/test_dedup.py @@ -380,3 +380,123 @@ def test_stats_aggregates_by_source(tmp_db: Path) -> None: assert stats.total == 3 assert stats.by_source == {"cls": 2, "sina": 1} assert stats.earliest is not None + + +# --------------------------------------------------------------------------- # +# 多源记录(source_ids) +# --------------------------------------------------------------------------- # + +def test_fingerprint_source_ids_default_to_source() -> None: + """source_ids 未显式给定时,自动包含主源 source_id。""" + fp = Fingerprint( + url_hash="h", content_hash="c", simhash=0, + source_id="cls", url="u", title="t", + ) + assert fp.source_ids == ["cls"] + + +def test_fingerprint_source_ids_keeps_main_source_first() -> None: + """source_ids 无论怎么传,主源 source_id 始终居首且去重。""" + fp = Fingerprint( + url_hash="h", content_hash="c", simhash=0, + source_id="cls", url="u", title="t", + source_ids=["sina", "cls", "eastmoney", "sina"], + ) + assert fp.source_ids[0] == "cls" + assert len(fp.source_ids) == len(set(fp.source_ids)) # 无重复 + + +def test_store_persists_source_ids(tmp_db: Path) -> None: + fp = Fingerprint( + url_hash="h", content_hash="c", simhash=0, + source_id="cls", url="u", title="t", + source_ids=["cls", "sina", "eastmoney"], + ) + with FingerprintStore(tmp_db) as store: + store.upsert(fp) + got = store.get_by_url_hash("h") + assert got is not None + assert got.source_ids == ["cls", "sina", "eastmoney"] + + +def test_store_migrates_old_schema_without_source_ids(tmp_db: Path) -> None: + """旧库(无 source_ids 列)打开时应自动迁移,旧数据回退为 [source_id]。""" + import sqlite3 + + conn = sqlite3.connect(tmp_db) + conn.executescript( + "CREATE TABLE fingerprints (" + " url_hash TEXT PRIMARY KEY, content_hash TEXT NOT NULL, simhash_hex TEXT NOT NULL," + " source_id TEXT NOT NULL, url TEXT NOT NULL, title TEXT NOT NULL," + " publish_date TEXT, ingested_at TEXT NOT NULL);" + ) + conn.execute( + "INSERT INTO fingerprints VALUES (?, ?, ?, ?, ?, ?, ?, ?)", + ("old1", "ch1", "0000000000000000", "cls", "u1", "t1", "2026-06-01", "2026-06-01T00:00:00"), + ) + conn.commit() + conn.close() + + with FingerprintStore(tmp_db) as store: + got = store.get_by_url_hash("old1") + assert got is not None + assert got.source_ids == ["cls"] # 迁移后回退主源 + # 迁移后可正常写入多源 + store.upsert(Fingerprint( + url_hash="new1", content_hash="c2", simhash=1, + source_id="sina", url="u2", title="t2", + source_ids=["sina", "cls"], + )) + assert store.get_by_url_hash("new1") is not None # type: ignore[union-attr] + + +def test_ingest_merges_sources_on_duplicate(tmp_db: Path) -> None: + """同一内容被多个源发布时,重复文章的来源并入唯一新闻指纹。""" + body = "宁德时代今日发布新一代麒麟电池,能量密度 255Wh/kg。" * 4 + a1 = _article(source_id="cls", url="https://cls/a", url_hash="aaaa111111111111", + content=body) + a2 = _article(source_id="sina", url="https://sina/b", url_hash="bbbb222222222222", + content=body) + a3 = _article(source_id="eastmoney", url="https://em/c", url_hash="cccc333333333333", + content=body) + + with Deduper(db_path=tmp_db) as d: + r1 = d.ingest(a1) + assert not r1.is_duplicate + r2 = d.ingest(a2) + assert r2.is_duplicate + assert r2.matched_layer == DedupLayer.CONTENT + assert r2.matched_source_id == "cls" + # 命中后 all_source_ids 立即包含两个源 + assert r2.all_source_ids == ["cls", "sina"] + + r3 = d.ingest(a3) + assert r3.is_duplicate + assert r3.all_source_ids == ["cls", "sina", "eastmoney"] + + # 指纹库持久化多源 + matched = d.store.get_by_url_hash("aaaa111111111111") + assert matched is not None + assert matched.source_ids == ["cls", "sina", "eastmoney"] + assert d.stats().total == 1 # 内容组只算 1 条唯一 + + +def test_check_reports_all_sources_without_writing(tmp_db: Path) -> None: + """check(只读)命中重复时也能看到全部来源,且不写库。""" + body = "宁德时代发布新一代麒麟电池产品。" * 6 + a1 = _article(source_id="cls", url="https://cls/a", url_hash="aaaa111111111111", + content=body) + a2 = _article(source_id="sina", url="https://sina/b", url_hash="bbbb222222222222", + content=body) + + with Deduper(db_path=tmp_db) as d: + d.ingest(a1) + d.ingest(a2) + # 第三次来一篇同样内容的文章,仅 check + a3 = _article(source_id="eastmoney", url="https://em/c", url_hash="cccc333333333333", + content=body) + result = d.check(a3) + assert result.is_duplicate + assert result.matched_source_id == "cls" + assert result.all_source_ids == ["cls", "sina"] + assert d.stats().total == 1 # check 不写库 diff --git a/tests/test_embedding.py b/tests/test_embedding.py index 6512d50..8385332 100644 --- a/tests/test_embedding.py +++ b/tests/test_embedding.py @@ -167,6 +167,27 @@ def test_resolve_provider_type_env_override(monkeypatch: pytest.MonkeyPatch) -> assert resolve_provider_type() == EmbeddingProviderType.LOCAL_BGE +def test_resolve_provider_type_scene_override(monkeypatch: pytest.MonkeyPatch) -> None: + """configs/llm_models.yaml 的 scenes.embedding.provider 优先于 .env。""" + monkeypatch.delenv("EMBEDDING_PROVIDER", raising=False) + monkeypatch.setattr( + "embedding.factory.load_scene_config", + lambda scene: {"provider": "local-bge"} if scene == "embedding" else {}, + ) + assert resolve_provider_type() == EmbeddingProviderType.LOCAL_BGE + + +def test_resolve_provider_type_explicit_arg_beats_scene( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """显式参数优先级最高,覆盖 YAML 场景。""" + monkeypatch.setattr( + "embedding.factory.load_scene_config", + lambda scene: {"provider": "local-bge"} if scene == "embedding" else {}, + ) + assert resolve_provider_type("dashscope") == EmbeddingProviderType.DASHSCOPE + + def test_resolve_provider_type_unknown_raises() -> None: with pytest.raises(EmbeddingError): resolve_provider_type("anthropic-emb") diff --git a/tests/test_llm.py b/tests/test_llm.py index 505e132..e128ffa 100644 --- a/tests/test_llm.py +++ b/tests/test_llm.py @@ -362,3 +362,114 @@ def test_load_llm_config_missing_model_raises(monkeypatch: pytest.MonkeyPatch) - monkeypatch.delenv("LLM_MODEL", raising=False) with pytest.raises(ValueError, match="模型"): load_llm_config(provider="deepseek") + + +# --------------------------------------------------------------------------- # +# load_llm_config —— configs/llm_models.yaml 场景配置 +# --------------------------------------------------------------------------- # + +def _patch_scene(monkeypatch: pytest.MonkeyPatch, cfg: dict) -> None: + """替换场景加载,模拟 configs/llm_models.yaml 中的某场景配置。""" + monkeypatch.setattr( + "llm.client.load_scene_config", + lambda scene: cfg if scene == "daily_report" else {}, + ) + + +def test_load_llm_config_scene_overrides_env(monkeypatch: pytest.MonkeyPatch) -> None: + """YAML 场景配置优先于 .env:provider / model / temperature / timeout / max_attempts。""" + monkeypatch.setenv("DEEPSEEK_API_KEY", "sk-env") + monkeypatch.setenv("DEEPSEEK_MODEL", "deepseek-env-model") + monkeypatch.setenv("QWEN_API_KEY", "sk-qwen") + _patch_scene(monkeypatch, { + "provider": "qwen", + "model": "qwen-max", + "temperature": 0.5, + "timeout_sec": 99, + "max_attempts": 5, + }) + cfg = load_llm_config(scene="daily_report") + assert cfg.provider == "qwen" + assert cfg.model == "qwen-max" + assert cfg.api_key == "sk-qwen" + assert cfg.temperature == 0.5 + assert cfg.timeout_sec == 99 + assert cfg.max_attempts == 5 + + +def test_load_llm_config_scene_blank_fields_fall_back_to_env( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """YAML 场景未配置的字段(如 model 留空)回退 .env,保持向后兼容。""" + monkeypatch.setenv("DEEPSEEK_API_KEY", "sk-env") + monkeypatch.setenv("DEEPSEEK_MODEL", "deepseek-env-model") + _patch_scene(monkeypatch, {"provider": "deepseek", "model": "", "temperature": 0.7}) + cfg = load_llm_config(scene="daily_report") + assert cfg.provider == "deepseek" + assert cfg.model == "deepseek-env-model" + assert cfg.api_key == "sk-env" + assert cfg.temperature == 0.7 + + +def test_load_llm_config_scene_api_key_env_name(monkeypatch: pytest.MonkeyPatch) -> None: + """api_key_env 指向自定义环境变量时,优先使用该变量。""" + monkeypatch.setenv("MY_CUSTOM_KEY", "sk-custom") + monkeypatch.setenv("DEEPSEEK_API_KEY", "sk-default") + monkeypatch.setenv("DEEPSEEK_MODEL", "deepseek-m") + _patch_scene(monkeypatch, { + "provider": "deepseek", + "model": "deepseek-scene-m", + "api_key_env": "MY_CUSTOM_KEY", + }) + cfg = load_llm_config(scene="daily_report") + assert cfg.api_key == "sk-custom" + assert cfg.model == "deepseek-scene-m" + + +def test_load_llm_config_scene_explicit_args_win(monkeypatch: pytest.MonkeyPatch) -> None: + """CLI/显式参数优先级最高,覆盖 YAML 场景。""" + monkeypatch.setenv("DEEPSEEK_API_KEY", "sk-env") + _patch_scene(monkeypatch, {"provider": "qwen", "model": "qwen-max"}) + monkeypatch.setenv("QWEN_API_KEY", "sk-qwen") + cfg = load_llm_config(provider="deepseek", model="deepseek-chat", scene="daily_report") + assert cfg.provider == "deepseek" + assert cfg.model == "deepseek-chat" + + +def test_load_llm_config_scene_missing_model_raises(monkeypatch: pytest.MonkeyPatch) -> None: + """场景与 .env 都未配置模型时必须报错(无内置兜底)。""" + monkeypatch.setenv("DEEPSEEK_API_KEY", "sk-test") + monkeypatch.delenv("DEEPSEEK_MODEL", raising=False) + monkeypatch.delenv("LLM_MODEL", raising=False) + _patch_scene(monkeypatch, {"provider": "deepseek", "model": ""}) + with pytest.raises(ValueError, match="模型"): + load_llm_config(scene="daily_report") + + +def test_load_llm_config_real_yaml_parseable() -> None: + """真实 configs/llm_models.yaml 必须可解析且包含全部场景(回归保护)。""" + from configs.loader import load_defaults, load_scene_config + + for scene in ("event_extraction", "daily_report", "stock_report", "embedding"): + assert isinstance(load_scene_config(scene), dict) + assert isinstance(load_defaults(), dict) + + +def test_load_llm_config_temperature_zero_is_respected( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """temperature=0 是合法配置,不应被 or 链回退默认。""" + monkeypatch.setenv("DEEPSEEK_API_KEY", "sk-env") + monkeypatch.setenv("DEEPSEEK_MODEL", "deepseek-m") + _patch_scene(monkeypatch, {"provider": "deepseek", "model": "deepseek-m", "temperature": 0}) + cfg = load_llm_config(scene="daily_report") + assert cfg.temperature == 0.0 + + +def test_load_llm_config_scene_max_attempts_zero() -> None: + """max_attempts=0 由 _pick_int 显式处理。""" + from llm.client import _pick_int + + assert _pick_int({"max_attempts": 0}, "max_attempts", 3) == 0 + assert _pick_int({"max_attempts": ""}, "max_attempts", 3) == 3 + assert _pick_int({}, "max_attempts", 3) == 3 diff --git a/tests/test_report_builder.py b/tests/test_report_builder.py index 3bd1d64..3aa844f 100644 --- a/tests/test_report_builder.py +++ b/tests/test_report_builder.py @@ -113,10 +113,20 @@ class TestLlmCallRetry: return SimpleNamespace(chat=SimpleNamespace(completions=Completions())), n + @staticmethod + def _cfg(): + from llm.client import LLMConfig + + return LLMConfig( + provider="deepseek", model="deepseek-v4-flash", + api_key="sk-test", base_url="https://api.deepseek.com", + temperature=0.3, + ) + def test_success_first_try(self) -> None: from scheduler.reporter import _llm_call client, n = self._fake_client(0) - out = _llm_call(client, "deepseek-v4-flash", "p") + out = _llm_call(client, self._cfg(), "p") assert out == "今日要点摘要" assert n["count"] == 1 @@ -125,7 +135,7 @@ class TestLlmCallRetry: monkeypatch.setattr(rep, "_LLM_RETRY_TIMES", 3) monkeypatch.setattr(rep, "_LLM_RETRY_BACKOFF_SEC", 0.01) client, n = self._fake_client(2) # 前 2 次失败,第 3 次成功 - out = rep._llm_call(client, "deepseek-v4-flash", "p") + out = rep._llm_call(client, self._cfg(), "p") assert out == "今日要点摘要" assert n["count"] == 3 @@ -135,7 +145,7 @@ class TestLlmCallRetry: monkeypatch.setattr(rep, "_LLM_RETRY_BACKOFF_SEC", 0.01) client, n = self._fake_client(99) # 一直失败 with pytest.raises(ConnectionError): - rep._llm_call(client, "deepseek-v4-flash", "p") + rep._llm_call(client, self._cfg(), "p") assert n["count"] == 2 # 重试 2 次后放弃