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:
+7
-1
@@ -8,15 +8,18 @@
|
||||
"""
|
||||
|
||||
from .client import (
|
||||
DEFAULT_MAX_ATTEMPTS,
|
||||
DEFAULT_TEMPERATURE,
|
||||
DEFAULT_TIMEOUT_SEC,
|
||||
SCENE_DAILY_REPORT,
|
||||
SCENE_EVENT_EXTRACTION,
|
||||
SCENE_STOCK_REPORT,
|
||||
LLMConfig,
|
||||
load_llm_config,
|
||||
make_async_client,
|
||||
make_sync_client,
|
||||
)
|
||||
from .extractor import (
|
||||
DEFAULT_MAX_ATTEMPTS,
|
||||
DEFAULT_PROMPT_PATH,
|
||||
MAX_CONTENT_CHARS,
|
||||
PromptTemplate,
|
||||
@@ -43,6 +46,9 @@ __all__ = [
|
||||
"MAX_CONTENT_CHARS",
|
||||
"MAX_IMPORTANCE",
|
||||
"MIN_IMPORTANCE",
|
||||
"SCENE_DAILY_REPORT",
|
||||
"SCENE_EVENT_EXTRACTION",
|
||||
"SCENE_STOCK_REPORT",
|
||||
"EventExtraction",
|
||||
"ExtractedEvent",
|
||||
"LLMCallError",
|
||||
|
||||
+108
-20
@@ -2,7 +2,13 @@
|
||||
|
||||
支持 DeepSeek 和 Qwen(百炼),两者均为 OpenAI 兼容接口,共用 openai SDK。
|
||||
|
||||
环境变量:
|
||||
配置来源(优先级从高到低):
|
||||
1. 代码 / CLI 显式参数(provider / model)
|
||||
2. configs/llm_models.yaml 场景配置(scene 参数,见 configs/loader.py)
|
||||
3. 环境变量 / .env(LLM_PROVIDER、DEEPSEEK_MODEL 等,向后兼容)
|
||||
4. 代码内置默认值
|
||||
|
||||
环境变量(兜底):
|
||||
LLM_PROVIDER = deepseek | qwen (默认 deepseek)
|
||||
DeepSeek: DEEPSEEK_API_KEY / DEEPSEEK_BASE_URL / DEEPSEEK_MODEL
|
||||
Qwen: QWEN_API_KEY / QWEN_BASE_URL / QWEN_MODEL
|
||||
@@ -19,6 +25,8 @@ from dataclasses import dataclass
|
||||
from loguru import logger
|
||||
from openai import AsyncOpenAI, OpenAI
|
||||
|
||||
from configs.loader import load_defaults, load_scene_config
|
||||
|
||||
# 默认基址
|
||||
_DEEPSEEK_DEFAULT_BASE = "https://api.deepseek.com"
|
||||
_QWEN_DEFAULT_BASE = "https://dashscope.aliyuncs.com/compatible-mode/v1"
|
||||
@@ -26,6 +34,12 @@ _QWEN_DEFAULT_BASE = "https://dashscope.aliyuncs.com/compatible-mode/v1"
|
||||
# 抽取任务默认参数
|
||||
DEFAULT_TIMEOUT_SEC = 60.0
|
||||
DEFAULT_TEMPERATURE = 0.1
|
||||
DEFAULT_MAX_ATTEMPTS = 3
|
||||
|
||||
# 场景名 -> configs/llm_models.yaml 中 scenes 的 key
|
||||
SCENE_EVENT_EXTRACTION = "event_extraction"
|
||||
SCENE_DAILY_REPORT = "daily_report"
|
||||
SCENE_STOCK_REPORT = "stock_report"
|
||||
|
||||
|
||||
@dataclass
|
||||
@@ -38,6 +52,7 @@ class LLMConfig:
|
||||
base_url: str
|
||||
timeout_sec: float = DEFAULT_TIMEOUT_SEC
|
||||
temperature: float = DEFAULT_TEMPERATURE
|
||||
max_attempts: int = DEFAULT_MAX_ATTEMPTS # 单次任务失败重试次数
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
if not self.api_key:
|
||||
@@ -51,49 +66,122 @@ def _read_env(key: str, default: str | None = None) -> str | None:
|
||||
return val.strip()
|
||||
|
||||
|
||||
def _first_env(keys: list[str | None]) -> str | None:
|
||||
"""按顺序返回第一个非空的环境变量值。"""
|
||||
for k in keys:
|
||||
if not k:
|
||||
continue
|
||||
v = _read_env(k)
|
||||
if v:
|
||||
return v
|
||||
return None
|
||||
|
||||
|
||||
def _num(value: object) -> float | None:
|
||||
"""把 YAML 数字/字符串安全转 float;非法或为空返回 None。"""
|
||||
if value is None or value == "":
|
||||
return None
|
||||
try:
|
||||
return float(value)
|
||||
except (TypeError, ValueError):
|
||||
return None
|
||||
|
||||
|
||||
def load_llm_config(
|
||||
provider: str | None = None,
|
||||
*,
|
||||
model: str | None = None,
|
||||
scene: str | None = None,
|
||||
) -> LLMConfig:
|
||||
"""根据环境变量构造 LLMConfig。
|
||||
"""按优先级构造 LLMConfig:显式参数 > YAML 场景 > 环境变量 > 内置默认。
|
||||
|
||||
provider 为 None 时读 LLM_PROVIDER 环境变量,默认 deepseek。
|
||||
model 为 None 时读 LLM_MODEL 或 provider 默认。
|
||||
scene 对应 configs/llm_models.yaml 中 scenes 的 key
|
||||
(event_extraction / daily_report / stock_report),该场景未配置的字段
|
||||
回退到环境变量,保持向后兼容。
|
||||
"""
|
||||
p = (provider or _read_env("LLM_PROVIDER", "deepseek") or "deepseek").lower()
|
||||
sc = load_scene_config(scene or "")
|
||||
dflt = load_defaults()
|
||||
|
||||
p = (
|
||||
provider
|
||||
or sc.get("provider")
|
||||
or _read_env("LLM_PROVIDER", "deepseek")
|
||||
or "deepseek"
|
||||
).lower()
|
||||
|
||||
# 各 provider 的 api_key / base_url / model 环境变量链
|
||||
provider_envs: dict[str, tuple[list[str | None], list[str | None], list[str | None]]] = {
|
||||
"deepseek": (
|
||||
[sc.get("api_key_env"), "DEEPSEEK_API_KEY"],
|
||||
[sc.get("base_url_env"), "DEEPSEEK_BASE_URL"],
|
||||
["DEEPSEEK_MODEL", "LLM_MODEL"],
|
||||
),
|
||||
"qwen": (
|
||||
[sc.get("api_key_env"), "QWEN_API_KEY", "DASHSCOPE_API_KEY"],
|
||||
[sc.get("base_url_env"), "QWEN_BASE_URL"],
|
||||
["QWEN_MODEL", "LLM_MODEL"],
|
||||
),
|
||||
}
|
||||
|
||||
if p == "deepseek":
|
||||
api_key = _read_env("DEEPSEEK_API_KEY") or ""
|
||||
base = _read_env("DEEPSEEK_BASE_URL", _DEEPSEEK_DEFAULT_BASE) or _DEEPSEEK_DEFAULT_BASE
|
||||
# DEEPSEEK_MODEL → LLM_MODEL;模型必须显式配置,不提供内置默认
|
||||
m = model or _read_env("DEEPSEEK_MODEL") or _read_env("LLM_MODEL")
|
||||
if not m:
|
||||
raise ValueError("未配置 LLM 模型: 请设置 DEEPSEEK_MODEL 或 LLM_MODEL")
|
||||
key_envs, base_envs, model_envs = provider_envs["deepseek"]
|
||||
default_base = _DEEPSEEK_DEFAULT_BASE
|
||||
elif p in ("qwen", "dashscope"):
|
||||
api_key = _read_env("QWEN_API_KEY") or _read_env("DASHSCOPE_API_KEY") or ""
|
||||
base = _read_env("QWEN_BASE_URL", _QWEN_DEFAULT_BASE) or _QWEN_DEFAULT_BASE
|
||||
# QWEN_MODEL → LLM_MODEL;模型必须显式配置,不提供内置默认
|
||||
m = model or _read_env("QWEN_MODEL") or _read_env("LLM_MODEL")
|
||||
if not m:
|
||||
raise ValueError("未配置 LLM 模型: 请设置 QWEN_MODEL 或 LLM_MODEL")
|
||||
key_envs, base_envs, model_envs = provider_envs["qwen"]
|
||||
default_base = _QWEN_DEFAULT_BASE
|
||||
p = "qwen" # 内部统一用 qwen
|
||||
else:
|
||||
raise ValueError(f"未知 LLM provider: {p!r},仅支持 deepseek / qwen")
|
||||
|
||||
timeout = float(_read_env("LLM_TIMEOUT_SEC", str(DEFAULT_TIMEOUT_SEC)) or DEFAULT_TIMEOUT_SEC)
|
||||
temperature = float(_read_env("LLM_TEMPERATURE", str(DEFAULT_TEMPERATURE)) or DEFAULT_TEMPERATURE)
|
||||
api_key = _first_env(key_envs) or ""
|
||||
base_url = _first_env(base_envs) or default_base
|
||||
# 模型优先级:显式参数 > YAML 场景 > 环境变量;模型必须显式配置,无内置兜底
|
||||
m = model or sc.get("model") or _first_env(model_envs)
|
||||
if not m:
|
||||
env_hint = "/".join(v for v in model_envs if v)
|
||||
raise ValueError(
|
||||
f"未配置 LLM 模型(场景 {scene or 'default'}): "
|
||||
f"请在 configs/llm_models.yaml 的 model 或 .env 设置 {env_hint}"
|
||||
)
|
||||
|
||||
timeout = _pick_float(sc, dflt, "timeout_sec", "LLM_TIMEOUT_SEC", DEFAULT_TIMEOUT_SEC)
|
||||
temperature = _pick_float(sc, dflt, "temperature", "LLM_TEMPERATURE", DEFAULT_TEMPERATURE)
|
||||
max_attempts = _pick_int(sc, "max_attempts", DEFAULT_MAX_ATTEMPTS)
|
||||
|
||||
return LLMConfig(
|
||||
provider=p,
|
||||
model=m,
|
||||
api_key=api_key,
|
||||
base_url=base,
|
||||
base_url=base_url,
|
||||
timeout_sec=timeout,
|
||||
temperature=temperature,
|
||||
max_attempts=max_attempts,
|
||||
)
|
||||
|
||||
|
||||
def _pick_float(
|
||||
sc: dict,
|
||||
dflt: dict,
|
||||
sc_key: str,
|
||||
env_key: str,
|
||||
default: float,
|
||||
) -> float:
|
||||
"""数值参数选择:YAML 场景 > 环境变量 > YAML defaults > 内置默认(零值合法)。"""
|
||||
v = _num(sc.get(sc_key))
|
||||
if v is not None:
|
||||
return v
|
||||
v = _num(_read_env(env_key))
|
||||
if v is not None:
|
||||
return v
|
||||
v = _num(dflt.get(sc_key))
|
||||
return v if v is not None else default
|
||||
|
||||
|
||||
def _pick_int(sc: dict, sc_key: str, default: int) -> int:
|
||||
v = _num(sc.get(sc_key))
|
||||
return int(v) if v is not None else default
|
||||
|
||||
|
||||
def make_sync_client(config: LLMConfig) -> OpenAI:
|
||||
"""构造同步 OpenAI 客户端(指向 DeepSeek/Qwen 兼容端点)。"""
|
||||
logger.debug(
|
||||
|
||||
+12
-4
@@ -187,11 +187,15 @@ def extract_event(
|
||||
article: Article,
|
||||
*,
|
||||
template: PromptTemplate | None = None,
|
||||
max_attempts: int = DEFAULT_MAX_ATTEMPTS,
|
||||
max_attempts: int | None = None,
|
||||
) -> ExtractedEvent:
|
||||
"""同步抽取单篇文章的事件(带重试)。"""
|
||||
"""同步抽取单篇文章的事件(带重试)。
|
||||
|
||||
max_attempts 为 None 时使用 config.max_attempts(来自 YAML/环境变量配置)。
|
||||
"""
|
||||
tpl = template or PromptTemplate()
|
||||
prompt = tpl.render(article)
|
||||
max_attempts = max_attempts or config.max_attempts
|
||||
|
||||
last_err: Exception | None = None
|
||||
for attempt in range(1, max_attempts + 1):
|
||||
@@ -241,12 +245,16 @@ async def extract_event_async(
|
||||
article: Article,
|
||||
*,
|
||||
template: PromptTemplate | None = None,
|
||||
max_attempts: int = DEFAULT_MAX_ATTEMPTS,
|
||||
max_attempts: int | None = None,
|
||||
semaphore: asyncio.Semaphore | None = None,
|
||||
) -> ExtractedEvent:
|
||||
"""异步抽取(批处理用),与同步版逻辑等价。"""
|
||||
"""异步抽取(批处理用),与同步版逻辑等价。
|
||||
|
||||
max_attempts 为 None 时使用 config.max_attempts。
|
||||
"""
|
||||
tpl = template or PromptTemplate()
|
||||
prompt = tpl.render(article)
|
||||
max_attempts = max_attempts or config.max_attempts
|
||||
|
||||
async def _run() -> ExtractedEvent:
|
||||
last_err: Exception | None = None
|
||||
|
||||
Reference in New Issue
Block a user