feat: 大模型使用场景化配置与去重多源记录

- 新增 configs/llm_models.yaml: 4 个场景(event_extraction/daily_report/stock_report/embedding)
  可独立配置 provider/model/api_key_env/base_url_env/temperature 等,含用途与模型要求说明
- 新增 configs/loader.py: YAML 场景加载器(优先级: CLI 参数 > YAML > .env > 内置默认)
- llm/client.py: load_llm_config 支持 scene 参数,LLMConfig 增加 max_attempts
- embedding/factory+remote+local: provider/model/batch_limit 支持场景覆盖
- scheduler/reporter+stock_reporter: 日报/个股摘要接入场景配置
- dedup: Fingerprint.source_ids 多源记录 + 旧库自动迁移 + DedupResult 多源字段
- scripts/run_dedup: uniques JSON 的 sources 字段 + data/deduped/{day}/sources.json 汇总
- scripts/run_event_extraction: 接入 event_extraction 场景
- 补充测试: 场景优先级/零值、多源合并、旧库迁移、embedding 场景覆盖
This commit is contained in:
2026-08-12 07:57:10 +08:00
parent 0c032196d2
commit 3c65701449
21 changed files with 886 additions and 80 deletions
+3
View File
@@ -34,6 +34,9 @@ htmlcov/
.DS_Store .DS_Store
Thumbs.db Thumbs.db
# 工具本地配置(非项目文件)
reasonix.toml
# 项目敏感配置 # 项目敏感配置
.env .env
.env.local .env.local
+1
View File
@@ -0,0 +1 @@
"""configs 配置包:提供 configs/*.yaml 的加载能力。"""
+133
View File
@@ -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,勿调大)
+69
View File
@@ -0,0 +1,69 @@
"""configs/ 目录下 YAML 配置加载器。
目前支持加载 configs/llm_models.yaml 的场景配置(scenes.<scene>)。
配置优先级(从高到低):
1. 代码 / 命令行显式参数(如 --provider qwen --model qwen-plus)
2. 本文件 YAML 场景配置(scenes.<scene>)
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.<scene> 的配置 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()
+32 -4
View File
@@ -35,7 +35,10 @@ def _publish_date(article: Article) -> str | None:
def article_to_fingerprint(article: Article) -> Fingerprint: def article_to_fingerprint(article: Article) -> Fingerprint:
"""构造 Fingerprint(用于 ingest 写入或对外只读)。""" """构造 Fingerprint(用于 ingest 写入或对外只读)。
source_ids 初始为 [article.source_id],后续重复文章命中时由 ingest 合并。
"""
return Fingerprint( return Fingerprint(
url_hash=article.url_hash, url_hash=article.url_hash,
content_hash=content_hash(article.content), content_hash=content_hash(article.content),
@@ -45,6 +48,7 @@ def article_to_fingerprint(article: Article) -> Fingerprint:
title=article.title, title=article.title,
publish_date=_publish_date(article), publish_date=_publish_date(article),
ingested_at=datetime.now(), ingested_at=datetime.now(),
source_ids=[article.source_id],
) )
@@ -80,7 +84,7 @@ class Deduper:
# ------------------------------------------------------------------ # # ------------------------------------------------------------------ #
def check(self, article: Article) -> DedupResult: def check(self, article: Article) -> DedupResult:
"""三层判重(只读)。""" """三层判重(只读)。命中时附带匹配指纹的多源信息(all_source_ids)。"""
fp = article_to_fingerprint(article) fp = article_to_fingerprint(article)
# L1: URL hash # L1: URL hash
@@ -93,6 +97,8 @@ class Deduper:
matched_url_hash=existing.url_hash, matched_url_hash=existing.url_hash,
matched_url=existing.url, matched_url=existing.url,
matched_title=existing.title, matched_title=existing.title,
matched_source_id=existing.source_id,
all_source_ids=existing.source_ids,
) )
# L2: 内容 hash # L2: 内容 hash
@@ -105,6 +111,8 @@ class Deduper:
matched_url_hash=existing.url_hash, matched_url_hash=existing.url_hash,
matched_url=existing.url, matched_url=existing.url,
matched_title=existing.title, matched_title=existing.title,
matched_source_id=existing.source_id,
all_source_ids=existing.source_ids,
) )
# L3: SimHash 模糊 # L3: SimHash 模糊
@@ -129,19 +137,39 @@ class Deduper:
matched_url_hash=best_match.url_hash, matched_url_hash=best_match.url_hash,
matched_url=best_match.url, matched_url=best_match.url,
matched_title=best_match.title, matched_title=best_match.title,
matched_source_id=best_match.source_id,
all_source_ids=best_match.source_ids,
hamming_distance=best_dist, hamming_distance=best_dist,
) )
return DedupResult(url_hash=fp.url_hash, is_duplicate=False) return DedupResult(url_hash=fp.url_hash, is_duplicate=False)
def ingest(self, article: Article) -> DedupResult: def ingest(self, article: Article) -> DedupResult:
"""判重 + 不重复则入库。""" """判重 + 不重复则入库。
命中重复时,把当前文章的 source_id 合并进匹配指纹的 source_ids
(记录同一内容组的全部来源),并更新 all_source_ids 后返回。
"""
result = self.check(article) result = self.check(article)
if not result.is_duplicate: if not result.is_duplicate:
fp = article_to_fingerprint(article) fp = article_to_fingerprint(article)
self.store.upsert(fp) self.store.upsert(fp)
logger.debug("入库: {} {}", fp.url_hash, fp.title[:30]) logger.debug("入库: {} {}", fp.url_hash, fp.title[:30])
else: 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()) logger.debug("命中重复: {}", result.short_summary())
return result return result
+29 -3
View File
@@ -4,9 +4,9 @@ from __future__ import annotations
from datetime import datetime from datetime import datetime
from enum import StrEnum 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): class DedupLayer(StrEnum):
@@ -18,7 +18,12 @@ class DedupLayer(StrEnum):
class Fingerprint(BaseModel): class Fingerprint(BaseModel):
"""单篇文章的指纹记录,持久化到 SQLite。""" """单篇文章的指纹记录,持久化到 SQLite。
source_ids: 同一内容组(去重后视为同一篇新闻)的全部来源列表,
第一位是主源(即本指纹的 source_id);重复文章命中时由
Deduper.ingest 自动合并,实现「一条唯一新闻记录多个源」。
"""
url_hash: str = Field(..., description="主键,与 Article.url_hash 一致") url_hash: str = Field(..., description="主键,与 Article.url_hash 一致")
content_hash: str = Field(..., description="标准化 content 的 SHA1[:16]") content_hash: str = Field(..., description="标准化 content 的 SHA1[:16]")
@@ -28,6 +33,20 @@ class Fingerprint(BaseModel):
title: str title: str
publish_date: str | None = Field(default=None, description="YYYY-MM-DD,用于时间窗口") publish_date: str | None = Field(default=None, description="YYYY-MM-DD,用于时间窗口")
ingested_at: datetime = Field(default_factory=datetime.now) 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): class DedupResult(BaseModel):
@@ -39,6 +58,13 @@ class DedupResult(BaseModel):
matched_url_hash: str | None = None matched_url_hash: str | None = None
matched_url: str | None = None matched_url: str | None = None
matched_title: 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( hamming_distance: int | None = Field(
default=None, description="仅 SimHash 层有值" default=None, description="仅 SimHash 层有值"
) )
+40 -2
View File
@@ -3,10 +3,14 @@
注意:SimHash 是 64 位无符号整数,SQLite INTEGER 是 64 位有符号 注意:SimHash 是 64 位无符号整数,SQLite INTEGER 是 64 位有符号
(范围 [-2^63, 2^63-1])。直接存可能溢出/转负数,虽然 XOR 仍然 (范围 [-2^63, 2^63-1])。直接存可能溢出/转负数,虽然 XOR 仍然
正确但语义混乱。这里统一存为 16 位 hex TEXT,避免符号问题。 正确但语义混乱。这里统一存为 16 位 hex TEXT,避免符号问题。
source_ids 列存 JSON 数组文本(同一内容组全部来源);旧库无此列时
自动 ALTER TABLE 迁移,旧数据读取时回退为 [source_id]。
""" """
from __future__ import annotations from __future__ import annotations
import json
import sqlite3 import sqlite3
from datetime import datetime, timedelta from datetime import datetime, timedelta
from pathlib import Path from pathlib import Path
@@ -27,13 +31,17 @@ CREATE TABLE IF NOT EXISTS fingerprints (
url TEXT NOT NULL, url TEXT NOT NULL,
title TEXT NOT NULL, title TEXT NOT NULL,
publish_date TEXT, 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_content_hash ON fingerprints(content_hash);
CREATE INDEX IF NOT EXISTS idx_publish_date ON fingerprints(publish_date); CREATE INDEX IF NOT EXISTS idx_publish_date ON fingerprints(publish_date);
CREATE INDEX IF NOT EXISTS idx_source_id ON fingerprints(source_id); 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: def _to_hex(simhash: int) -> str:
return f"{simhash:016x}" return f"{simhash:016x}"
@@ -43,6 +51,25 @@ def _from_hex(hex_str: str) -> int:
return int(hex_str, 16) 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: def _row_to_fp(row: sqlite3.Row) -> Fingerprint:
return Fingerprint( return Fingerprint(
url_hash=row["url_hash"], url_hash=row["url_hash"],
@@ -53,6 +80,7 @@ def _row_to_fp(row: sqlite3.Row) -> Fingerprint:
title=row["title"], title=row["title"],
publish_date=row["publish_date"], publish_date=row["publish_date"],
ingested_at=datetime.fromisoformat(row["ingested_at"]), 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.row_factory = sqlite3.Row
self._conn.executescript(_SCHEMA_SQL) self._conn.executescript(_SCHEMA_SQL)
self._migrate_source_ids()
logger.debug("打开指纹库: {}", self.db_path) 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: def close(self) -> None:
self._conn.close() self._conn.close()
@@ -148,7 +185,7 @@ class FingerprintStore:
self._conn.execute( self._conn.execute(
"INSERT OR REPLACE INTO fingerprints " "INSERT OR REPLACE INTO fingerprints "
"(url_hash, content_hash, simhash_hex, source_id, url, title, " "(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.url_hash,
fp.content_hash, fp.content_hash,
@@ -158,6 +195,7 @@ class FingerprintStore:
fp.title, fp.title,
fp.publish_date, fp.publish_date,
fp.ingested_at.isoformat(), fp.ingested_at.isoformat(),
_to_sources_json(fp.source_ids),
), ),
) )
+14 -3
View File
@@ -1,9 +1,14 @@
"""Embedding provider 工厂:根据环境变量构造合适后端。""" """Embedding provider 工厂:根据配置构造合适后端。
配置优先级: 显式参数 > configs/llm_models.yaml scenes.embedding > .env > 默认。
"""
from __future__ import annotations from __future__ import annotations
import os import os
from configs.loader import load_scene_config
from .base import AsyncEmbeddingProvider, EmbeddingProvider from .base import AsyncEmbeddingProvider, EmbeddingProvider
from .models import EmbeddingError, EmbeddingProviderType from .models import EmbeddingError, EmbeddingProviderType
from .remote import ( 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: def resolve_provider_type(provider: str | None = None) -> EmbeddingProviderType:
"""根据 provider 参数 / env 解析出 EmbeddingProviderType。 """根据 provider 参数 / YAML 场景 / env 解析出 EmbeddingProviderType。
映射: 映射:
dashscope / qwen / remote -> DASHSCOPE dashscope / qwen / remote -> DASHSCOPE
local / local-bge / bge / bge-m3 -> LOCAL_BGE local / local-bge / bge / bge-m3 -> LOCAL_BGE
默认 dashscope。 默认 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"): if p in ("dashscope", "qwen", "remote"):
return EmbeddingProviderType.DASHSCOPE return EmbeddingProviderType.DASHSCOPE
if p in ("local", "local-bge", "bge", "bge-m3"): if p in ("local", "local-bge", "bge", "bge-m3"):
+5 -1
View File
@@ -21,6 +21,8 @@ from .models import EmbeddingError
if TYPE_CHECKING: if TYPE_CHECKING:
from sentence_transformers import SentenceTransformer from sentence_transformers import SentenceTransformer
from configs.loader import load_scene_config
LOCAL_DEFAULT_MODEL = "BAAI/bge-m3" LOCAL_DEFAULT_MODEL = "BAAI/bge-m3"
LOCAL_DEFAULT_DIM = 1024 LOCAL_DEFAULT_DIM = 1024
@@ -56,10 +58,12 @@ class LocalBGEEmbeddingProvider(EmbeddingProvider):
normalize: bool = True, normalize: bool = True,
) -> None: ) -> None:
st_cls = _try_import_st() st_cls = _try_import_st()
# 模型优先级:CLI/参数 > LOCAL_EMBEDDING_MODEL > 默认值 # 模型优先级:CLI/参数 > YAML scenes.embedding > LOCAL_EMBEDDING_MODEL > 默认值
# 不读全局 EMBEDDING_MODEL,避免与 DashScope 冲突 # 不读全局 EMBEDDING_MODEL,避免与 DashScope 冲突
scene_model = load_scene_config("embedding").get("model")
self.model = ( self.model = (
model model
or scene_model
or _read_env("LOCAL_EMBEDDING_MODEL", LOCAL_DEFAULT_MODEL) or _read_env("LOCAL_EMBEDDING_MODEL", LOCAL_DEFAULT_MODEL)
or LOCAL_DEFAULT_MODEL or LOCAL_DEFAULT_MODEL
) )
+59 -19
View File
@@ -5,10 +5,14 @@
model: text-embedding-v3 (1024 维) model: text-embedding-v3 (1024 维)
限制: 单次请求 input ≤ 25 条 限制: 单次请求 input ≤ 25 条
环境变量: 配置来源(优先级从高到低):
DASHSCOPE_API_KEY 1. 构造参数(model / api_key / base_url / max_attempts)
QWEN_BASE_URL (默认百炼兼容路径) 2. configs/llm_models.yaml 的 scenes.embedding
EMBEDDING_MODEL (默认 text-embedding-v3) 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 from __future__ import annotations
@@ -19,6 +23,8 @@ import os
from loguru import logger from loguru import logger
from openai import AsyncOpenAI, OpenAI from openai import AsyncOpenAI, OpenAI
from configs.loader import load_scene_config
from .base import AsyncEmbeddingProvider, EmbeddingProvider from .base import AsyncEmbeddingProvider, EmbeddingProvider
from .models import EmbeddingError from .models import EmbeddingError
@@ -32,6 +38,9 @@ DEFAULT_MAX_ATTEMPTS = 3
RETRY_BASE_WAIT_SEC = 1.0 RETRY_BASE_WAIT_SEC = 1.0
RETRY_MAX_WAIT_SEC = 8.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: def _read_env(key: str, default: str | None = None) -> str | None:
val = os.environ.get(key) val = os.environ.get(key)
@@ -40,23 +49,48 @@ def _read_env(key: str, default: str | None = None) -> str | None:
return val.strip() 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]: def _resolve_config() -> tuple[str, str, str]:
"""读取 API key / base_url / model,返回 (api_key, base_url, model)。 """读取 API key / base_url / model,返回 (api_key, base_url, model)。
优先级: YAML 场景 > 环境变量 > 内置默认。
模型名优先级:DASHSCOPE_EMBEDDING_MODEL > 默认值。 模型名优先级:DASHSCOPE_EMBEDDING_MODEL > 默认值。
不再读全局 EMBEDDING_MODEL,避免与 LOCAL provider 冲突。 不再读全局 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: if not api_key:
raise EmbeddingError("DASHSCOPE_EMBEDDING_API_KEY 或 DASHSCOPE_API_KEY 未配置") raise EmbeddingError(
# DASHSCOPE_EMBEDDING_BASE_URL -> QWEN_BASE_URL(兜底) -> 默认 f"{api_key_env} 或 DASHSCOPE_API_KEY 未配置"
)
# YAML base_url_env -> DASHSCOPE_EMBEDDING_BASE_URL -> QWEN_BASE_URL(兜底) -> 默认
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 _read_env("QWEN_BASE_URL")
or DASHSCOPE_DEFAULT_BASE or DASHSCOPE_DEFAULT_BASE
) )
model = ( 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 or DASHSCOPE_DEFAULT_MODEL
) )
return api_key, base_url, model return api_key, base_url, model
@@ -78,24 +112,27 @@ class DashScopeEmbeddingProvider(EmbeddingProvider):
model: str | None = None, model: str | None = None,
api_key: str | None = None, api_key: str | None = None,
base_url: str | None = None, base_url: str | None = None,
timeout_sec: float = 60.0, timeout_sec: float | None = None,
max_attempts: int = DEFAULT_MAX_ATTEMPTS, max_attempts: int | None = None,
batch_limit: int | None = None,
) -> None: ) -> None:
env_key, env_base, env_model = _resolve_config() env_key, env_base, env_model = _resolve_config()
self.model = model or env_model self.model = model or env_model
self.dim = DASHSCOPE_DEFAULT_DIM 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( self._client = OpenAI(
api_key=api_key or env_key, api_key=api_key or env_key,
base_url=base_url or env_base, base_url=base_url or env_base,
timeout=timeout_sec, timeout=timeout,
) )
def embed_batch(self, texts: list[str]) -> list[list[float]]: def embed_batch(self, texts: list[str]) -> list[list[float]]:
if not texts: if not texts:
return [] return []
results: list[list[float]] = [] 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)) results.extend(self._call_with_retry(chunk))
return results return results
@@ -136,24 +173,27 @@ class DashScopeAsyncEmbeddingProvider(AsyncEmbeddingProvider):
model: str | None = None, model: str | None = None,
api_key: str | None = None, api_key: str | None = None,
base_url: str | None = None, base_url: str | None = None,
timeout_sec: float = 60.0, timeout_sec: float | None = None,
max_attempts: int = DEFAULT_MAX_ATTEMPTS, max_attempts: int | None = None,
batch_limit: int | None = None,
) -> None: ) -> None:
env_key, env_base, env_model = _resolve_config() env_key, env_base, env_model = _resolve_config()
self.model = model or env_model self.model = model or env_model
self.dim = DASHSCOPE_DEFAULT_DIM 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( self._client = AsyncOpenAI(
api_key=api_key or env_key, api_key=api_key or env_key,
base_url=base_url or env_base, base_url=base_url or env_base,
timeout=timeout_sec, timeout=timeout,
) )
async def embed_batch(self, texts: list[str]) -> list[list[float]]: async def embed_batch(self, texts: list[str]) -> list[list[float]]:
if not texts: if not texts:
return [] return []
results: list[list[float]] = [] 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)) results.extend(await self._call_with_retry(chunk))
return results return results
+7 -1
View File
@@ -8,15 +8,18 @@
""" """
from .client import ( from .client import (
DEFAULT_MAX_ATTEMPTS,
DEFAULT_TEMPERATURE, DEFAULT_TEMPERATURE,
DEFAULT_TIMEOUT_SEC, DEFAULT_TIMEOUT_SEC,
SCENE_DAILY_REPORT,
SCENE_EVENT_EXTRACTION,
SCENE_STOCK_REPORT,
LLMConfig, LLMConfig,
load_llm_config, load_llm_config,
make_async_client, make_async_client,
make_sync_client, make_sync_client,
) )
from .extractor import ( from .extractor import (
DEFAULT_MAX_ATTEMPTS,
DEFAULT_PROMPT_PATH, DEFAULT_PROMPT_PATH,
MAX_CONTENT_CHARS, MAX_CONTENT_CHARS,
PromptTemplate, PromptTemplate,
@@ -43,6 +46,9 @@ __all__ = [
"MAX_CONTENT_CHARS", "MAX_CONTENT_CHARS",
"MAX_IMPORTANCE", "MAX_IMPORTANCE",
"MIN_IMPORTANCE", "MIN_IMPORTANCE",
"SCENE_DAILY_REPORT",
"SCENE_EVENT_EXTRACTION",
"SCENE_STOCK_REPORT",
"EventExtraction", "EventExtraction",
"ExtractedEvent", "ExtractedEvent",
"LLMCallError", "LLMCallError",
+108 -20
View File
@@ -2,7 +2,13 @@
支持 DeepSeek 和 Qwen(百炼),两者均为 OpenAI 兼容接口,共用 openai SDK。 支持 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) LLM_PROVIDER = deepseek | qwen (默认 deepseek)
DeepSeek: DEEPSEEK_API_KEY / DEEPSEEK_BASE_URL / DEEPSEEK_MODEL DeepSeek: DEEPSEEK_API_KEY / DEEPSEEK_BASE_URL / DEEPSEEK_MODEL
Qwen: QWEN_API_KEY / QWEN_BASE_URL / QWEN_MODEL Qwen: QWEN_API_KEY / QWEN_BASE_URL / QWEN_MODEL
@@ -19,6 +25,8 @@ from dataclasses import dataclass
from loguru import logger from loguru import logger
from openai import AsyncOpenAI, OpenAI from openai import AsyncOpenAI, OpenAI
from configs.loader import load_defaults, load_scene_config
# 默认基址 # 默认基址
_DEEPSEEK_DEFAULT_BASE = "https://api.deepseek.com" _DEEPSEEK_DEFAULT_BASE = "https://api.deepseek.com"
_QWEN_DEFAULT_BASE = "https://dashscope.aliyuncs.com/compatible-mode/v1" _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_TIMEOUT_SEC = 60.0
DEFAULT_TEMPERATURE = 0.1 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 @dataclass
@@ -38,6 +52,7 @@ class LLMConfig:
base_url: str base_url: str
timeout_sec: float = DEFAULT_TIMEOUT_SEC timeout_sec: float = DEFAULT_TIMEOUT_SEC
temperature: float = DEFAULT_TEMPERATURE temperature: float = DEFAULT_TEMPERATURE
max_attempts: int = DEFAULT_MAX_ATTEMPTS # 单次任务失败重试次数
def __post_init__(self) -> None: def __post_init__(self) -> None:
if not self.api_key: if not self.api_key:
@@ -51,49 +66,122 @@ def _read_env(key: str, default: str | None = None) -> str | None:
return val.strip() 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( def load_llm_config(
provider: str | None = None, provider: str | None = None,
*, *,
model: str | None = None, model: str | None = None,
scene: str | None = None,
) -> LLMConfig: ) -> LLMConfig:
"""根据环境变量构造 LLMConfig。 """按优先级构造 LLMConfig:显式参数 > YAML 场景 > 环境变量 > 内置默认。
provider 为 None 时读 LLM_PROVIDER 环境变量,默认 deepseek。 scene 对应 configs/llm_models.yaml 中 scenes 的 key
model 为 None 时读 LLM_MODEL 或 provider 默认。 (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": if p == "deepseek":
api_key = _read_env("DEEPSEEK_API_KEY") or "" key_envs, base_envs, model_envs = provider_envs["deepseek"]
base = _read_env("DEEPSEEK_BASE_URL", _DEEPSEEK_DEFAULT_BASE) or _DEEPSEEK_DEFAULT_BASE default_base = _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")
elif p in ("qwen", "dashscope"): elif p in ("qwen", "dashscope"):
api_key = _read_env("QWEN_API_KEY") or _read_env("DASHSCOPE_API_KEY") or "" key_envs, base_envs, model_envs = provider_envs["qwen"]
base = _read_env("QWEN_BASE_URL", _QWEN_DEFAULT_BASE) or _QWEN_DEFAULT_BASE default_base = _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")
p = "qwen" # 内部统一用 qwen p = "qwen" # 内部统一用 qwen
else: else:
raise ValueError(f"未知 LLM provider: {p!r},仅支持 deepseek / qwen") raise ValueError(f"未知 LLM provider: {p!r},仅支持 deepseek / qwen")
timeout = float(_read_env("LLM_TIMEOUT_SEC", str(DEFAULT_TIMEOUT_SEC)) or DEFAULT_TIMEOUT_SEC) api_key = _first_env(key_envs) or ""
temperature = float(_read_env("LLM_TEMPERATURE", str(DEFAULT_TEMPERATURE)) or DEFAULT_TEMPERATURE) 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( return LLMConfig(
provider=p, provider=p,
model=m, model=m,
api_key=api_key, api_key=api_key,
base_url=base, base_url=base_url,
timeout_sec=timeout, timeout_sec=timeout,
temperature=temperature, 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: def make_sync_client(config: LLMConfig) -> OpenAI:
"""构造同步 OpenAI 客户端(指向 DeepSeek/Qwen 兼容端点)。""" """构造同步 OpenAI 客户端(指向 DeepSeek/Qwen 兼容端点)。"""
logger.debug( logger.debug(
+12 -4
View File
@@ -187,11 +187,15 @@ def extract_event(
article: Article, article: Article,
*, *,
template: PromptTemplate | None = None, template: PromptTemplate | None = None,
max_attempts: int = DEFAULT_MAX_ATTEMPTS, max_attempts: int | None = None,
) -> ExtractedEvent: ) -> ExtractedEvent:
"""同步抽取单篇文章的事件(带重试)。""" """同步抽取单篇文章的事件(带重试)。
max_attempts 为 None 时使用 config.max_attempts(来自 YAML/环境变量配置)。
"""
tpl = template or PromptTemplate() tpl = template or PromptTemplate()
prompt = tpl.render(article) prompt = tpl.render(article)
max_attempts = max_attempts or config.max_attempts
last_err: Exception | None = None last_err: Exception | None = None
for attempt in range(1, max_attempts + 1): for attempt in range(1, max_attempts + 1):
@@ -241,12 +245,16 @@ async def extract_event_async(
article: Article, article: Article,
*, *,
template: PromptTemplate | None = None, template: PromptTemplate | None = None,
max_attempts: int = DEFAULT_MAX_ATTEMPTS, max_attempts: int | None = None,
semaphore: asyncio.Semaphore | None = None, semaphore: asyncio.Semaphore | None = None,
) -> ExtractedEvent: ) -> ExtractedEvent:
"""异步抽取(批处理用),与同步版逻辑等价。""" """异步抽取(批处理用),与同步版逻辑等价。
max_attempts 为 None 时使用 config.max_attempts。
"""
tpl = template or PromptTemplate() tpl = template or PromptTemplate()
prompt = tpl.render(article) prompt = tpl.render(article)
max_attempts = max_attempts or config.max_attempts
async def _run() -> ExtractedEvent: async def _run() -> ExtractedEvent:
last_err: Exception | None = None last_err: Exception | None = None
+17 -9
View File
@@ -17,7 +17,10 @@ import time
from collections import Counter from collections import Counter
from datetime import date, datetime, timedelta from datetime import date, datetime, timedelta
from pathlib import Path 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 dotenv import load_dotenv
from loguru import logger from loguru import logger
@@ -519,10 +522,10 @@ def _generate_ai_summary(news: dict, cninfo: dict, day_str: str,
return "" return ""
try: try:
from llm.client import load_llm_config, make_sync_client from llm.client import SCENE_DAILY_REPORT, load_llm_config, make_sync_client
config = load_llm_config() config = load_llm_config(scene=SCENE_DAILY_REPORT)
client = make_sync_client(config) 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: except Exception as e:
logger.warning("AI 摘要生成失败: {}", e) logger.warning("AI 摘要生成失败: {}", e)
return "" return ""
@@ -548,8 +551,12 @@ def _split_lines_into_chunks(lines: list[str], max_chars: int = 3000) -> list[li
return chunks return chunks
def _llm_summarize(client, model: str, lines: list[str], day_str: str) -> str: def _llm_summarize(client, config: LLMConfig, lines: list[str], day_str: str) -> str:
"""LLM 摘要:单块直接总结,多块先分段总结再合并。""" """LLM 摘要:单块直接总结,多块先分段总结再合并。
config 为 llm.client.LLMConfig(daily_report 场景),提供 model / temperature。
"""
model = config.model
chunks = _split_lines_into_chunks(lines) chunks = _split_lines_into_chunks(lines)
if len(chunks) == 1: 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 后的文本。 """单次 LLM 调用(带重试),返回 strip 后的文本。
config 为 llm.client.LLMConfig(daily_report 场景),提供 model / temperature。
失败按指数退避重试 `_LLM_RETRY_TIMES` 次(默认 3),全部失败则抛出最后一次异常。 失败按指数退避重试 `_LLM_RETRY_TIMES` 次(默认 3),全部失败则抛出最后一次异常。
若 finish_reason 为 'length' 则说明达到 max_tokens 上限被截断。 若 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): for attempt in range(_LLM_RETRY_TIMES):
try: try:
resp = client.chat.completions.create( resp = client.chat.completions.create(
model=model, model=config.model,
messages=[ messages=[
{"role": "system", "content": "你是 A 股日报撰写助手,输出简洁、有洞察的新闻摘要。"}, {"role": "system", "content": "你是 A 股日报撰写助手,输出简洁、有洞察的新闻摘要。"},
{"role": "user", "content": prompt}, {"role": "user", "content": prompt},
], ],
temperature=0.3, temperature=config.temperature,
max_tokens=max_tokens, max_tokens=max_tokens,
) )
content = (resp.choices[0].message.content or "").strip() content = (resp.choices[0].message.content or "").strip()
+3 -3
View File
@@ -232,7 +232,7 @@ def _generate_ai_summary(company_name: str, announcements: list[dict],
news: list[dict], research: list[dict], news: list[dict], research: list[dict],
irm: list[dict]) -> str: irm: list[dict]) -> str:
"""LLM 生成个股要点分析。""" """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 = [] lines = []
@@ -277,12 +277,12 @@ def _generate_ai_summary(company_name: str, announcements: list[dict],
直接输出要点列表:""" 直接输出要点列表:"""
try: try:
config = load_llm_config() config = load_llm_config(scene=SCENE_STOCK_REPORT)
client = make_sync_client(config) client = make_sync_client(config)
resp = client.chat.completions.create( resp = client.chat.completions.create(
model=config.model, model=config.model,
messages=[{"role": "user", "content": prompt}], 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() return (resp.choices[0].message.content or "").strip()
except Exception as e: except Exception as e:
+83 -6
View File
@@ -2,9 +2,12 @@
输入: data/processed/{source}/{YYYYMMDD}/*.json (M2 产物) 输入: data/processed/{source}/{YYYYMMDD}/*.json (M2 产物)
输出: 输出:
- 指纹库:data/dedup/fingerprints.sqlite3 - 指纹库:data/dedup/fingerprints.sqlite3 (source_ids 列记录多源)
- 唯一文章:data/deduped/{YYYYMMDD}/uniques/{url_hash}.json - 唯一文章:data/deduped/{YYYYMMDD}/uniques/{url_hash}.json (含 sources 多源字段)
- 多源记录:data/deduped/{YYYYMMDD}/sources.json
{url_hash: [source_id, ...]},一条唯一新闻的全部来源
- 重复记录:data/deduped/{YYYYMMDD}/duplicates.jsonl - 重复记录:data/deduped/{YYYYMMDD}/duplicates.jsonl
(含 matched_source_id / matched_source_ids)
用法: 用法:
uv run python -m scripts.run_dedup # 处理今日全部源 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()) 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( def _process_source_day(
source_id: str, source_id: str,
day: str, day: str,
processed_root: Path, processed_root: Path,
out_root: Path, out_root: Path,
deduper: Deduper, deduper: Deduper,
sources_map: dict[str, list[str]],
) -> tuple[int, int, Counter]: ) -> 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 src_dir = processed_root / source_id / day
if not src_dir.is_dir(): if not src_dir.is_dir():
logger.info("源 {} 日期 {} 无 processed 目录,跳过", source_id, day) logger.info("源 {} 日期 {} 无 processed 目录,跳过", source_id, day)
@@ -92,6 +141,11 @@ def _process_source_day(
dup_cnt += 1 dup_cnt += 1
if result.matched_layer is not None: if result.matched_layer is not None:
layer_cnt[result.matched_layer.value] += 1 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( dup_f.write(
json.dumps( json.dumps(
{ {
@@ -107,6 +161,8 @@ def _process_source_day(
"matched_url": result.matched_url, "matched_url": result.matched_url,
"matched_url_hash": result.matched_url_hash, "matched_url_hash": result.matched_url_hash,
"matched_title": result.matched_title, "matched_title": result.matched_title,
"matched_source_id": result.matched_source_id,
"matched_source_ids": result.all_source_ids,
"hamming_distance": result.hamming_distance, "hamming_distance": result.hamming_distance,
}, },
ensure_ascii=False, ensure_ascii=False,
@@ -115,8 +171,8 @@ def _process_source_day(
) )
else: else:
uniq_cnt += 1 uniq_cnt += 1
out_path = uniques_dir / f"{article.url_hash}.json" sources_map[article.url_hash] = [article.source_id]
out_path.write_text(article.model_dump_json(indent=2), encoding="utf-8") _write_unique(article.url_hash, article, uniques_dir, sources_map)
total = uniq_cnt + dup_cnt total = uniq_cnt + dup_cnt
rate = dup_cnt / max(total, 1) rate = dup_cnt / max(total, 1)
@@ -173,14 +229,35 @@ def main() -> int:
total_uniq = 0 total_uniq = 0
total_dup = 0 total_dup = 0
total_layers: Counter = Counter() total_layers: Counter = Counter()
# 当天唯一新闻 url_hash -> 全部来源列表(跨源累积,多源记录)
sources_map: dict[str, list[str]] = {}
for src in sources: for src in sources:
u, d, lc = _process_source_day( 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_uniq += u
total_dup += d total_dup += d
total_layers.update(lc) 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 total = total_uniq + total_dup
rate = total_dup / max(total, 1) rate = total_dup / max(total, 1)
logger.info( logger.info(
+5 -1
View File
@@ -31,6 +31,7 @@ from pydantic import ValidationError
from extractor import Article from extractor import Article
from llm import ( from llm import (
SCENE_EVENT_EXTRACTION,
ExtractedEvent, ExtractedEvent,
LLMCallError, LLMCallError,
PromptTemplate, PromptTemplate,
@@ -88,7 +89,10 @@ def _load_article(p: Path) -> Article | 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
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( logger.info(
"LLM provider={} model={} base_url={}", "LLM provider={} model={} base_url={}",
config.provider, config.model, config.base_url, config.provider, config.model, config.base_url,
+120
View File
@@ -380,3 +380,123 @@ def test_stats_aggregates_by_source(tmp_db: Path) -> None:
assert stats.total == 3 assert stats.total == 3
assert stats.by_source == {"cls": 2, "sina": 1} assert stats.by_source == {"cls": 2, "sina": 1}
assert stats.earliest is not None 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 不写库
+21
View File
@@ -167,6 +167,27 @@ def test_resolve_provider_type_env_override(monkeypatch: pytest.MonkeyPatch) ->
assert resolve_provider_type() == EmbeddingProviderType.LOCAL_BGE 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: def test_resolve_provider_type_unknown_raises() -> None:
with pytest.raises(EmbeddingError): with pytest.raises(EmbeddingError):
resolve_provider_type("anthropic-emb") resolve_provider_type("anthropic-emb")
+111
View File
@@ -362,3 +362,114 @@ def test_load_llm_config_missing_model_raises(monkeypatch: pytest.MonkeyPatch) -
monkeypatch.delenv("LLM_MODEL", raising=False) monkeypatch.delenv("LLM_MODEL", raising=False)
with pytest.raises(ValueError, match="模型"): with pytest.raises(ValueError, match="模型"):
load_llm_config(provider="deepseek") 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
+13 -3
View File
@@ -113,10 +113,20 @@ class TestLlmCallRetry:
return SimpleNamespace(chat=SimpleNamespace(completions=Completions())), n 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: def test_success_first_try(self) -> None:
from scheduler.reporter import _llm_call from scheduler.reporter import _llm_call
client, n = self._fake_client(0) client, n = self._fake_client(0)
out = _llm_call(client, "deepseek-v4-flash", "p") out = _llm_call(client, self._cfg(), "p")
assert out == "今日要点摘要" assert out == "今日要点摘要"
assert n["count"] == 1 assert n["count"] == 1
@@ -125,7 +135,7 @@ class TestLlmCallRetry:
monkeypatch.setattr(rep, "_LLM_RETRY_TIMES", 3) monkeypatch.setattr(rep, "_LLM_RETRY_TIMES", 3)
monkeypatch.setattr(rep, "_LLM_RETRY_BACKOFF_SEC", 0.01) monkeypatch.setattr(rep, "_LLM_RETRY_BACKOFF_SEC", 0.01)
client, n = self._fake_client(2) # 前 2 次失败,第 3 次成功 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 out == "今日要点摘要"
assert n["count"] == 3 assert n["count"] == 3
@@ -135,7 +145,7 @@ class TestLlmCallRetry:
monkeypatch.setattr(rep, "_LLM_RETRY_BACKOFF_SEC", 0.01) monkeypatch.setattr(rep, "_LLM_RETRY_BACKOFF_SEC", 0.01)
client, n = self._fake_client(99) # 一直失败 client, n = self._fake_client(99) # 一直失败
with pytest.raises(ConnectionError): with pytest.raises(ConnectionError):
rep._llm_call(client, "deepseek-v4-flash", "p") rep._llm_call(client, self._cfg(), "p")
assert n["count"] == 2 # 重试 2 次后放弃 assert n["count"] == 2 # 重试 2 次后放弃