- 删除 5 个过时/残留文档(project_plan/agent_prompt/optimization_plan/report_db_design/deploy/README) - 新建 docs/architecture.md(项目架构:11 包职责+数据模型+配置+产物) - 重写 docs/user-guide.md(CLI 全量+增量/断点续跑+MCP+FAQ) - 重写 README.md(精简入口+文档索引) - 更新 continuation.md(追加本次记录) - 更新 .gitignore(排除 data/* 运行产物)
209 lines
6.5 KiB
Python
209 lines
6.5 KiB
Python
"""LLM 客户端抽象与工厂。
|
|
|
|
支持 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
|
|
(QWEN_API_KEY -> DASHSCOPE_API_KEY 兜底)
|
|
模型必须显式配置(provider 对应的 *_MODEL 或 LLM_MODEL),不再提供内置默认模型。
|
|
LLM_TEMPERATURE / LLM_TIMEOUT_SEC
|
|
"""
|
|
|
|
from __future__ import annotations
|
|
|
|
import os
|
|
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"
|
|
|
|
# 抽取任务默认参数
|
|
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
|
|
class LLMConfig:
|
|
"""LLM 调用配置(provider / model / api_key / base_url / 参数)。"""
|
|
|
|
provider: str # "deepseek" / "qwen"
|
|
model: str
|
|
api_key: str
|
|
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:
|
|
raise ValueError(f"LLM provider={self.provider} 的 API key 为空")
|
|
|
|
|
|
def _read_env(key: str, default: str | None = None) -> str | None:
|
|
val = os.environ.get(key)
|
|
if val is None or val.strip() == "":
|
|
return default
|
|
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:显式参数 > YAML 场景 > 环境变量 > 内置默认。
|
|
|
|
scene 对应 configs/llm_models.yaml 中 scenes 的 key
|
|
(event_extraction / daily_report / stock_report),该场景未配置的字段
|
|
回退到环境变量,保持向后兼容。
|
|
"""
|
|
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":
|
|
key_envs, base_envs, model_envs = provider_envs["deepseek"]
|
|
default_base = _DEEPSEEK_DEFAULT_BASE
|
|
elif p in ("qwen", "dashscope"):
|
|
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")
|
|
|
|
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_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(
|
|
"初始化同步 LLM 客户端: provider={} model={} base_url={}",
|
|
config.provider, config.model, config.base_url,
|
|
)
|
|
return OpenAI(
|
|
api_key=config.api_key,
|
|
base_url=config.base_url,
|
|
timeout=config.timeout_sec,
|
|
)
|
|
|
|
|
|
def make_async_client(config: LLMConfig) -> AsyncOpenAI:
|
|
"""构造异步 OpenAI 客户端(用于批处理高并发)。"""
|
|
logger.debug(
|
|
"初始化异步 LLM 客户端: provider={} model={} base_url={}",
|
|
config.provider, config.model, config.base_url,
|
|
)
|
|
return AsyncOpenAI(
|
|
api_key=config.api_key,
|
|
base_url=config.base_url,
|
|
timeout=config.timeout_sec,
|
|
)
|