Files
news/llm/client.py
T
2026-07-18 15:51:01 +08:00

120 lines
3.9 KiB
Python

"""LLM 客户端抽象与工厂。
支持 DeepSeek 和 Qwen(百炼),两者均为 OpenAI 兼容接口,共用 openai SDK。
环境变量:
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 兜底)
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
# 默认基址
_DEEPSEEK_DEFAULT_BASE = "https://api.deepseek.com"
_QWEN_DEFAULT_BASE = "https://dashscope.aliyuncs.com/compatible-mode/v1"
# 默认模型
_DEEPSEEK_DEFAULT_MODEL = "deepseek-chat"
_QWEN_DEFAULT_MODEL = "qwen-plus"
# 抽取任务默认参数
DEFAULT_TIMEOUT_SEC = 60.0
DEFAULT_TEMPERATURE = 0.1
@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
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 load_llm_config(
provider: str | None = None,
*,
model: str | None = None,
) -> LLMConfig:
"""根据环境变量构造 LLMConfig。
provider 为 None 时读 LLM_PROVIDER 环境变量,默认 deepseek。
model 为 None 时读 LLM_MODEL 或 provider 默认。
"""
p = (provider or _read_env("LLM_PROVIDER", "deepseek") or "deepseek").lower()
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") or _DEEPSEEK_DEFAULT_MODEL
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") or _QWEN_DEFAULT_MODEL
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)
return LLMConfig(
provider=p,
model=m,
api_key=api_key,
base_url=base,
timeout_sec=timeout,
temperature=temperature,
)
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,
)