Compare commits
2
Commits
0c032196d2
...
c1a803968a
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
c1a803968a | ||
|
|
3c65701449 |
@@ -5,6 +5,7 @@
|
|||||||
# ============================
|
# ============================
|
||||||
|
|
||||||
# ---- LLM Provider ----
|
# ---- LLM Provider ----
|
||||||
|
# 推荐: 各场景按需独立配置见 configs/llm_models.yaml(优先级高于以下环境变量)
|
||||||
LLM_PROVIDER=deepseek
|
LLM_PROVIDER=deepseek
|
||||||
|
|
||||||
# ---- DeepSeek ----
|
# ---- DeepSeek ----
|
||||||
|
|||||||
@@ -34,6 +34,9 @@ htmlcov/
|
|||||||
.DS_Store
|
.DS_Store
|
||||||
Thumbs.db
|
Thumbs.db
|
||||||
|
|
||||||
|
# 工具本地配置(非项目文件)
|
||||||
|
reasonix.toml
|
||||||
|
|
||||||
# 项目敏感配置
|
# 项目敏感配置
|
||||||
.env
|
.env
|
||||||
.env.local
|
.env.local
|
||||||
|
|||||||
@@ -209,12 +209,34 @@ uv run python -m scripts.run_dedup --reset
|
|||||||
|
|
||||||
去重产物路径:
|
去重产物路径:
|
||||||
|
|
||||||
- `data/dedup/fingerprints.sqlite3` 指纹库(跨日累积)
|
- `data/dedup/fingerprints.sqlite3` 指纹库(跨日累积,`source_ids` 列记录多源)
|
||||||
- `data/deduped/{YYYYMMDD}/uniques/{url_hash}.json` 唯一文章(可送 M4+ 处理)
|
- `data/deduped/{YYYYMMDD}/uniques/{url_hash}.json` 唯一文章(可送 M4+ 处理,含 `sources` 多源字段)
|
||||||
- `data/deduped/{YYYYMMDD}/duplicates.jsonl` 重复记录(含命中层 / 命中目标)
|
- `data/deduped/{YYYYMMDD}/sources.json` 多源记录 `{url_hash: [source_id, ...]}`
|
||||||
|
- `data/deduped/{YYYYMMDD}/duplicates.jsonl` 重复记录(含命中层 / 命中目标 / matched_source_ids)
|
||||||
|
|
||||||
|
**多源记录**:同一内容被多个新闻源发布时,去重后只保留一条唯一新闻,但会记录全部来源
|
||||||
|
(主源居首)。指纹库 `source_ids` 列为跨日累积的权威记录;`uniques/{url_hash}.json` 的
|
||||||
|
`sources` 字段与 `data/deduped/{YYYYMMDD}/sources.json`(本次去重涉及内容组的源汇总,
|
||||||
|
含跨日命中)为当日产物,展示/下游可按需读取多源列表。
|
||||||
|
|
||||||
日志:`logs/dedup.log`
|
日志:`logs/dedup.log`
|
||||||
|
|
||||||
|
### 大模型使用场景配置(configs/llm_models.yaml)
|
||||||
|
|
||||||
|
项目共 4 个调用大模型的场景,均可独立配置 provider / model,见 `configs/llm_models.yaml`
|
||||||
|
(文件内有每个场景的用途、使用方式、对模型的要求的完整说明):
|
||||||
|
|
||||||
|
| 场景 key | 用途 | 调用方 | 模型要求 |
|
||||||
|
| --- | --- | --- | --- |
|
||||||
|
| `event_extraction` | 投资事件抽取(输出严格 JSON) | M4 `llm/extractor.py` | OpenAI 兼容 + JSON mode + 上下文 ≥16K |
|
||||||
|
| `daily_report` | 日报 AI 摘要(500 字要点) | `scheduler/reporter.py` | 长输出(≥1500 tokens)、中文摘要 |
|
||||||
|
| `stock_report` | 个股 AI 要点分析 | `scheduler/stock_reporter.py` | 长输出(≥500 tokens)、要点化 |
|
||||||
|
| `embedding` | 文本向量化(入库/检索) | M5、MCP、pipeline | 嵌入模型,1024 维,dashscope / local-bge |
|
||||||
|
|
||||||
|
配置优先级:**CLI 显式参数 > YAML 场景 > .env 环境变量 > 代码默认**。
|
||||||
|
YAML 中字段留空即回退 `.env`(向后兼容,现有部署无需改动即可运行)。
|
||||||
|
API Key 一律放 `.env`,YAML 只保存环境变量名(`api_key_env`)。
|
||||||
|
|
||||||
### 运行 LLM 事件抽取(M4)
|
### 运行 LLM 事件抽取(M4)
|
||||||
|
|
||||||
M4 调用 DeepSeek 或 Qwen,从 M3 唯一文章中抽取结构化投资事件(stock_codes/sentiment/importance/event_type 等)。
|
M4 调用 DeepSeek 或 Qwen,从 M3 唯一文章中抽取结构化投资事件(stock_codes/sentiment/importance/event_type 等)。
|
||||||
@@ -253,6 +275,9 @@ uv run python -m scripts.run_event_extraction --no-deduped --source cls
|
|||||||
- `DEEPSEEK_API_KEY` / `DEEPSEEK_BASE_URL`
|
- `DEEPSEEK_API_KEY` / `DEEPSEEK_BASE_URL`
|
||||||
- `DASHSCOPE_API_KEY` / `QWEN_BASE_URL`
|
- `DASHSCOPE_API_KEY` / `QWEN_BASE_URL`
|
||||||
|
|
||||||
|
> 注:以上环境变量作为兜底;按场景独立配置 provider/model 推荐使用
|
||||||
|
> `configs/llm_models.yaml`(见上文「大模型使用场景配置」)。
|
||||||
|
|
||||||
日志:`logs/llm.log`
|
日志:`logs/llm.log`
|
||||||
|
|
||||||
### 运行 Embedding 向量化(M5)
|
### 运行 Embedding 向量化(M5)
|
||||||
@@ -288,6 +313,9 @@ uv run python -m scripts.run_embedding --provider local-bge
|
|||||||
- `LOCAL_EMBEDDING_MODEL` 本地模型(默认 `BAAI/bge-m3`)
|
- `LOCAL_EMBEDDING_MODEL` 本地模型(默认 `BAAI/bge-m3`)
|
||||||
- `DASHSCOPE_API_KEY` / `QWEN_BASE_URL`(M4 已配)
|
- `DASHSCOPE_API_KEY` / `QWEN_BASE_URL`(M4 已配)
|
||||||
|
|
||||||
|
> 注:以上环境变量作为兜底;embedding 场景的 provider/model 也可在
|
||||||
|
> `configs/llm_models.yaml` 的 `scenes.embedding` 中配置(优先级更高)。
|
||||||
|
|
||||||
日志:`logs/embedding.log`
|
日志:`logs/embedding.log`
|
||||||
|
|
||||||
### 运行 Qdrant 入库与检索(M6)
|
### 运行 Qdrant 入库与检索(M6)
|
||||||
|
|||||||
@@ -0,0 +1 @@
|
|||||||
|
"""configs 配置包:提供 configs/*.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,勿调大)
|
||||||
@@ -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()
|
||||||
+39
-2
@@ -1,6 +1,6 @@
|
|||||||
# continuation.md
|
# continuation.md
|
||||||
|
|
||||||
> `checkpoint` @ 2026-08-05 08:30
|
> `checkpoint` @ 2026-08-11 17:00
|
||||||
|
|
||||||
---
|
---
|
||||||
|
|
||||||
@@ -13,12 +13,49 @@
|
|||||||
| 日报 | **M10 完成并已部署 pi5: 结构化入库 MySQL;日报按当天日期生成(新闻 30h 回溯 / xwlb 取前一日 / 公告调研近 15 日)** |
|
| 日报 | **M10 完成并已部署 pi5: 结构化入库 MySQL;日报按当天日期生成(新闻 30h 回溯 / xwlb 取前一日 / 公告调研近 15 日)** |
|
||||||
| DB 连接 | pi 上 systemd 服务 `a-share-db-tunnel` 常驻(0.0.0.0:13306 → doorcome.cn:3306);**pi5 直连 192.168.1.10:13306** |
|
| DB 连接 | pi 上 systemd 服务 `a-share-db-tunnel` 常驻(0.0.0.0:13306 → doorcome.cn:3306);**pi5 直连 192.168.1.10:13306** |
|
||||||
| 调度器 | APScheduler,systemd `a-share-research.service`(pi5);每天 07:00 首次任务生成日报(12/18/22 点不生成) |
|
| 调度器 | APScheduler,systemd `a-share-research.service`(pi5);每天 07:00 首次任务生成日报(12/18/22 点不生成) |
|
||||||
| LLM | `deepseek-v4-flash`(绝不允许擅自修改;模型必须显式配置,无内置兜底) |
|
| LLM | 场景化配置 `configs/llm_models.yaml`(4 场景: event_extraction/daily_report/stock_report/embedding);YAML 优先、`.env` 兜底;模型必须显式配置,无内置兜底 |
|
||||||
|
| 去重 | 多源记录:指纹库 `source_ids` 列 + uniques JSON `sources` 字段 + `data/deduped/{day}/sources.json` |
|
||||||
| 服务器 | `pi@192.168.1.160`(生产)/ `pi@192.168.1.10`(DB 隧道宿主) |
|
| 服务器 | `pi@192.168.1.160`(生产)/ `pi@192.168.1.10`(DB 隧道宿主) |
|
||||||
| 抓取方式 | js_render=false → httpx 直连;js_render=true → Playwright |
|
| 抓取方式 | js_render=false → httpx 直连;js_render=true → Playwright |
|
||||||
|
|
||||||
---
|
---
|
||||||
|
|
||||||
|
## 本次完成 (2026-08-11) — 大模型场景化配置 + 去重多源记录
|
||||||
|
|
||||||
|
**目标:** ① 梳理全部 AI 大模型使用点,新增 `configs/llm_models.yaml` 按场景独立配置 provider/model;② 去重时记录一条唯一新闻的全部来源。
|
||||||
|
|
||||||
|
**1. 大模型使用点梳理(共 4 个场景,详见 configs/llm_models.yaml 内注释):**
|
||||||
|
- `event_extraction`(M4 投资事件抽取,JSON mode,llm/extractor.py)
|
||||||
|
- `daily_report`(日报 AI 摘要,scheduler/reporter.py)
|
||||||
|
- `stock_report`(个股 AI 要点分析,scheduler/stock_reporter.py)
|
||||||
|
- `embedding`(向量化,dashscope 远程 / local-bge 本地,embedding/remote.py + local.py)
|
||||||
|
- 非使用点确认:crawler 纯抓取、MCP 仅复用 embedding、run_xwlb 抓外部「AI 精编」数据源
|
||||||
|
|
||||||
|
**2. 场景配置实现(优先级: CLI 显式参数 > YAML > .env > 内置默认):**
|
||||||
|
- 新增 `configs/loader.py`(lru_cache 读 llm_models.yaml)+ `configs/__init__.py`
|
||||||
|
- `llm/client.py`:`load_llm_config(scene=...)` 支持场景;`LLMConfig` 增加 `max_attempts`;模型缺失仍报错(无内置兜底)
|
||||||
|
- `llm/extractor.py`:`extract_event(_async)` 的 max_attempts 默认取 `config.max_attempts`
|
||||||
|
- `embedding/factory.py` + `remote.py` + `local.py`:provider/model/api_key_env/base_url_env/batch_limit 支持场景覆盖
|
||||||
|
- `scheduler/reporter.py`(daily_report)+ `stock_reporter.py`(stock_report):接入场景,temperature 取配置
|
||||||
|
- YAML 中 provider/model 默认留空 → 回退 .env,**现有部署零改动兼容**
|
||||||
|
|
||||||
|
**3. 去重多源记录:**
|
||||||
|
- `dedup/models.py`:`Fingerprint.source_ids`(validator 保主源居首+去重);`DedupResult` 增 `matched_source_id`/`all_source_ids`
|
||||||
|
- `dedup/store.py`:指纹库加 `source_ids` 列,旧库自动 ALTER 迁移,旧数据回退 `[source_id]`
|
||||||
|
- `dedup/deduper.py`:`ingest` 命中重复时把新源合并进匹配指纹
|
||||||
|
- `scripts/run_dedup.py`:uniques JSON 附加 `sources` 字段;重复命中时仅更新 sources 不覆盖原文;输出 `data/deduped/{day}/sources.json` 汇总
|
||||||
|
|
||||||
|
**验证(Mac 本地):**
|
||||||
|
- 全量 pytest:**215 passed**(仅 crawler 3 个 retry mock 失败为基线预存在问题,与本改动无关)
|
||||||
|
- 端到端人工构造 3 源同文:1 条唯一 + sources.json `["cls","eastmoney","sina"]` + 指纹库 source_ids 列正确
|
||||||
|
- ruff:9 个错误均为基线既有(crawler/cninfo.py 未用 import、reporter.py L5/L4 命名),本次零新增
|
||||||
|
|
||||||
|
**待办:**
|
||||||
|
- 生产同步:代码 + `configs/llm_models.yaml` scp 到 pi5(注意 rsync 排除规则含 configs/*.yaml,需显式同步),重启 `a-share-research` 生效
|
||||||
|
- 首次同步前 pi5 无 YAML → 全部回退 .env,行为不变,可平滑切换
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
## 本次完成 (2026-08-05) — 日报可靠性修复与取数逻辑优化
|
## 本次完成 (2026-08-05) — 日报可靠性修复与取数逻辑优化
|
||||||
|
|
||||||
**目标:** 解决日报 AI 摘要偶发失败;修正日报日期与 xwlb/新闻取数语义。
|
**目标:** 解决日报 AI 摘要偶发失败;修正日报日期与 xwlb/新闻取数语义。
|
||||||
|
|||||||
+33
-5
@@ -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,20 +137,40 @@ 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
|
||||||
logger.debug("命中重复: {}", result.short_summary())
|
|
||||||
|
# 重复:合并来源到匹配指纹(主源保持首位,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
|
return result
|
||||||
|
|
||||||
def stats(self) -> DedupStats:
|
def stats(self) -> DedupStats:
|
||||||
|
|||||||
+29
-3
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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()
|
||||||
|
|||||||
@@ -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
@@ -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(
|
||||||
|
|||||||
@@ -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,
|
||||||
|
|||||||
@@ -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 不写库
|
||||||
|
|||||||
@@ -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")
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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 次后放弃
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user