feat(llm): 多模型 fallback 支持
Some checks failed
CI / pytest (Python 3.10) (push) Has been cancelled
CI / pytest (Python 3.11) (push) Has been cancelled
CI / pytest (Python 3.12) (push) Has been cancelled

LLMClient 改为供应商链模式:
- 主供应商 + GAOKAO_LLM_FALLBACK_MODELS/PROVIDERS/API_KEYS/BASE_URLS
- 按顺序尝试, 第一个成功即返回
- 全部失败时抛出聚合错误(含供应商数量和最后错误)

配置层:
- Settings 新增 4 个 fallback 字段
- load_settings() 读取对应环境变量

新增测试(4 个):
- test_fallback_config_parsed: 验证 fallback 链解析正确
- test_fallback_no_config: 无 fallback 时只有主供应商
- test_fallback_falls_through_to_second_provider: 主失败自动切换
- test_fallback_all_fail_raises: 全部失败抛聚合错误

部署示例:
.env.docker.example / .env.payment.example 补充注释掉的 fallback 模板

验证: data/llm/tests 16 passed, admin/tests 40 passed
This commit is contained in:
Hermes Agent
2026-06-28 13:34:17 +08:00
parent a3d8c73a84
commit 11fbb591c3
5 changed files with 214 additions and 12 deletions

View File

@@ -16,6 +16,11 @@ GAOKAO_LLM_BASE_URL=https://dashscope.aliyuncs.com/compatible-mode/v1
GAOKAO_LLM_MODEL=qwen-plus GAOKAO_LLM_MODEL=qwen-plus
GAOKAO_LLM_TIMEOUT=60 GAOKAO_LLM_TIMEOUT=60
GAOKAO_LLM_MAX_TOKENS=4096 GAOKAO_LLM_MAX_TOKENS=4096
# 可选:多模型 fallback逗号分隔主供应商失败时按顺序尝试
# GAOKAO_LLM_FALLBACK_MODELS=gpt-4o-mini,deepseek-chat
# GAOKAO_LLM_FALLBACK_PROVIDERS=openai,deepseek
# GAOKAO_LLM_FALLBACK_API_KEYS=sk-openai-key,sk-deepseek-key
# GAOKAO_LLM_FALLBACK_BASE_URLS=https://api.openai.com/v1,https://api.deepseek.com/v1
GAOKAO_PAYMENT_PROVIDER=mock GAOKAO_PAYMENT_PROVIDER=mock
GAOKAO_PAYMENT_BASE_URL=https://example.com GAOKAO_PAYMENT_BASE_URL=https://example.com
GAOKAO_PAYMENT_WEBHOOK_SECRET=replace-with-independent-payment-webhook-secret GAOKAO_PAYMENT_WEBHOOK_SECRET=replace-with-independent-payment-webhook-secret

View File

@@ -12,6 +12,11 @@ GAOKAO_LLM_BASE_URL=https://dashscope.aliyuncs.com/compatible-mode/v1
GAOKAO_LLM_MODEL=qwen-plus GAOKAO_LLM_MODEL=qwen-plus
GAOKAO_LLM_TIMEOUT=60 GAOKAO_LLM_TIMEOUT=60
GAOKAO_LLM_MAX_TOKENS=4096 GAOKAO_LLM_MAX_TOKENS=4096
# 可选:多模型 fallback逗号分隔主供应商失败时按顺序尝试
# GAOKAO_LLM_FALLBACK_MODELS=gpt-4o-mini,deepseek-chat
# GAOKAO_LLM_FALLBACK_PROVIDERS=openai,deepseek
# GAOKAO_LLM_FALLBACK_API_KEYS=sk-openai-key,sk-deepseek-key
# GAOKAO_LLM_FALLBACK_BASE_URLS=https://api.openai.com/v1,https://api.deepseek.com/v1
# ===== 支付/安全/运营 ===== # ===== 支付/安全/运营 =====
GAOKAO_PAYMENT_APP_ID= GAOKAO_PAYMENT_APP_ID=

View File

@@ -79,6 +79,13 @@ class Settings:
llm_model: str llm_model: str
llm_timeout_seconds: int llm_timeout_seconds: int
llm_max_tokens: int llm_max_tokens: int
# 多模型 fallback逗号分隔的额外模型按顺序尝试
llm_fallback_models: str # 例: "qwen-turbo,gpt-4o-mini,deepseek-chat"
llm_fallback_providers: (
str # 例: "dashscope,openai,deepseek"(与 fallback_models 对应)
)
llm_fallback_api_keys: str # 例: "key1,key2,key3"(与 fallback_models 对应)
llm_fallback_base_urls: str # 例: "url1,url2,url3"(与 fallback_models 对应)
def _resolve_payment_webhook_secret(env: str) -> str: def _resolve_payment_webhook_secret(env: str) -> str:
@@ -332,6 +339,10 @@ def load_settings() -> Settings:
llm_model=os.getenv("GAOKAO_LLM_MODEL", "qwen-plus"), llm_model=os.getenv("GAOKAO_LLM_MODEL", "qwen-plus"),
llm_timeout_seconds=int(os.getenv("GAOKAO_LLM_TIMEOUT", "60")), llm_timeout_seconds=int(os.getenv("GAOKAO_LLM_TIMEOUT", "60")),
llm_max_tokens=int(os.getenv("GAOKAO_LLM_MAX_TOKENS", "4096")), llm_max_tokens=int(os.getenv("GAOKAO_LLM_MAX_TOKENS", "4096")),
llm_fallback_models=os.getenv("GAOKAO_LLM_FALLBACK_MODELS", ""),
llm_fallback_providers=os.getenv("GAOKAO_LLM_FALLBACK_PROVIDERS", ""),
llm_fallback_api_keys=os.getenv("GAOKAO_LLM_FALLBACK_API_KEYS", ""),
llm_fallback_base_urls=os.getenv("GAOKAO_LLM_FALLBACK_BASE_URLS", ""),
) )
# 生产环境 post-load 校验:webhook / portal token / JWT / admin password # 生产环境 post-load 校验:webhook / portal token / JWT / admin password
# / payment provider 必须满足强度门槛, 任一不满足 fail-closed (P0-2/P2-4/P2-5/6/20)。 # / payment provider 必须满足强度门槛, 任一不满足 fail-closed (P0-2/P2-4/P2-5/6/20)。

View File

@@ -36,17 +36,84 @@ class LLMClient:
def __init__(self, settings: Settings) -> None: def __init__(self, settings: Settings) -> None:
self._settings = settings self._settings = settings
self._provider = settings.llm_provider
self._api_key = settings.llm_api_key
self._base_url = settings.llm_base_url.rstrip("/")
self._model = settings.llm_model
self._timeout = settings.llm_timeout_seconds self._timeout = settings.llm_timeout_seconds
self._max_tokens = settings.llm_max_tokens self._max_tokens = settings.llm_max_tokens
# 构建供应商链:主供应商 + fallback 供应商列表
self._providers: list[dict[str, str]] = []
if settings.llm_provider != "none" and settings.llm_api_key:
self._providers.append({
"provider": settings.llm_provider,
"api_key": settings.llm_api_key,
"base_url": settings.llm_base_url.rstrip("/"),
"model": settings.llm_model,
})
# 解析 fallback 配置
fb_models = [
s.strip()
for s in (settings.llm_fallback_models or "").split(",")
if s.strip()
]
fb_providers = [
s.strip()
for s in (settings.llm_fallback_providers or "").split(",")
if s.strip()
]
fb_keys = [
s.strip()
for s in (settings.llm_fallback_api_keys or "").split(",")
if s.strip()
]
fb_urls = [
s.strip()
for s in (settings.llm_fallback_base_urls or "").split(",")
if s.strip()
]
for i, model in enumerate(fb_models):
provider = (
fb_providers[i]
if i < len(fb_providers)
else (self._providers[0]["provider"] if self._providers else "openai")
)
api_key = (
fb_keys[i]
if i < len(fb_keys)
else (self._providers[0]["api_key"] if self._providers else "")
)
base_url = (
fb_urls[i].rstrip("/")
if i < len(fb_urls)
else (
self._providers[0]["base_url"]
if self._providers
else "https://api.openai.com/v1"
)
)
if api_key:
self._providers.append({
"provider": provider,
"api_key": api_key,
"base_url": base_url,
"model": model,
})
# 兼容旧接口
self._provider = self._providers[0]["provider"] if self._providers else "none"
self._api_key = self._providers[0]["api_key"] if self._providers else ""
self._base_url = self._providers[0]["base_url"] if self._providers else ""
self._model = self._providers[0]["model"] if self._providers else ""
@property @property
def is_configured(self) -> bool: def is_configured(self) -> bool:
"""LLM 是否已配置可用。""" """LLM 是否已配置可用。"""
return self._provider != "none" and bool(self._api_key) return len(self._providers) > 0
@property
def provider_count(self) -> int:
"""已配置的供应商数量(含主+fallback"""
return len(self._providers)
def chat( def chat(
self, self,
@@ -55,7 +122,10 @@ class LLMClient:
temperature: float = 0.7, temperature: float = 0.7,
max_tokens: int | None = None, max_tokens: int | None = None,
) -> LLMResponse: ) -> LLMResponse:
"""调用 chat completions API。 """调用 chat completions API,支持多供应商 fallback
按供应商链顺序依次尝试,第一个成功即返回。
全部失败时抛出最后一个错误。
Args: Args:
messages: OpenAI 格式的消息列表。 messages: OpenAI 格式的消息列表。
@@ -66,7 +136,7 @@ class LLMClient:
LLMResponse。 LLMResponse。
Raises: Raises:
LLMError: 调用失败。 LLMError: 全部供应商都失败。
""" """
if not self.is_configured: if not self.is_configured:
raise LLMError( raise LLMError(
@@ -74,17 +144,52 @@ class LLMClient:
"请设置 GAOKAO_LLM_PROVIDER 和 GAOKAO_LLM_API_KEY。" "请设置 GAOKAO_LLM_PROVIDER 和 GAOKAO_LLM_API_KEY。"
) )
last_error: LLMError | None = None
for idx, prov in enumerate(self._providers):
try:
return self._call_single_provider(
provider=prov,
messages=messages,
temperature=temperature,
max_tokens=max_tokens or self._max_tokens,
)
except LLMError as e:
last_error = e
prov_name = prov["provider"]
model_name = prov["model"]
# 如果还有下一个供应商,继续尝试
if idx < len(self._providers) - 1:
continue
# 最后一个也失败了
raise LLMError(
f"全部 {len(self._providers)} 个 LLM 供应商均失败。"
f"最后错误 ({prov_name}/{model_name}): {e}"
) from e
# 理论上不会到达这里
raise last_error or LLMError("未知 LLM 错误")
def _call_single_provider(
self,
*,
provider: dict[str, str],
messages: list[dict[str, str]],
temperature: float,
max_tokens: int,
) -> LLMResponse:
"""调用单个供应商的 API。"""
payload: dict[str, Any] = { payload: dict[str, Any] = {
"model": self._model, "model": provider["model"],
"messages": messages, "messages": messages,
"temperature": temperature, "temperature": temperature,
"max_tokens": max_tokens or self._max_tokens, "max_tokens": max_tokens,
} }
url = f"{self._base_url}/chat/completions" url = f"{provider['base_url']}/chat/completions"
headers = { headers = {
"Content-Type": "application/json", "Content-Type": "application/json",
"Authorization": f"Bearer {self._api_key}", "Authorization": f"Bearer {provider['api_key']}",
} }
data = json.dumps(payload).encode("utf-8") data = json.dumps(payload).encode("utf-8")
@@ -116,7 +221,7 @@ class LLMClient:
raise LLMError(f"LLM API 返回空 content: {result}") raise LLMError(f"LLM API 返回空 content: {result}")
usage = result.get("usage", {}) usage = result.get("usage", {})
model = result.get("model", self._model) model = result.get("model", provider["model"])
return LLMResponse( return LLMResponse(
content=content, content=content,

View File

@@ -22,6 +22,10 @@ class MockSettings:
llm_model: str = "qwen-plus" llm_model: str = "qwen-plus"
llm_timeout_seconds: int = 60 llm_timeout_seconds: int = 60
llm_max_tokens: int = 4096 llm_max_tokens: int = 4096
llm_fallback_models: str = ""
llm_fallback_providers: str = ""
llm_fallback_api_keys: str = ""
llm_fallback_base_urls: str = ""
class TestLLMClient: class TestLLMClient:
@@ -160,3 +164,75 @@ class TestPrompts:
assert "家长希望省内优先" in user assert "家长希望省内优先" in user
assert "volunteers" in user assert "volunteers" in user
assert "至少 8 条" in user assert "至少 8 条" in user
class TestFallback:
def test_fallback_config_parsed(self):
"""fallback 配置被正确解析为供应商链。"""
settings = MockSettings(
llm_provider="dashscope",
llm_api_key="key-main",
llm_model="qwen-plus",
llm_fallback_models="gpt-4o-mini,deepseek-chat",
llm_fallback_providers="openai,deepseek",
llm_fallback_api_keys="key-openai,key-deepseek",
llm_fallback_base_urls="https://api.openai.com/v1,https://api.deepseek.com/v1",
)
client = LLMClient(settings)
assert client.provider_count == 3
assert client._providers[0]["model"] == "qwen-plus"
assert client._providers[1]["model"] == "gpt-4o-mini"
assert client._providers[2]["model"] == "deepseek-chat"
assert client._providers[1]["api_key"] == "key-openai"
assert client._providers[2]["provider"] == "deepseek"
def test_fallback_no_config(self):
"""无 fallback 配置时只有主供应商。"""
client = LLMClient(MockSettings(llm_provider="openai", llm_api_key="sk-test"))
assert client.provider_count == 1
def test_fallback_falls_through_to_second_provider(self):
"""主供应商失败时,自动尝试第二个供应商。"""
settings = MockSettings(
llm_provider="dashscope",
llm_api_key="key-main",
llm_model="qwen-plus",
llm_fallback_models="gpt-4o-mini",
llm_fallback_providers="openai",
llm_fallback_api_keys="key-openai",
llm_fallback_base_urls="https://api.openai.com/v1",
)
client = LLMClient(settings)
call_count = {"n": 0}
def mock_call_side_effect(**kwargs):
call_count["n"] += 1
if call_count["n"] == 1:
raise LLMError("first provider failed")
return LLMResponse(content="fallback success", model="gpt-4o-mini")
with patch.object(
client, "_call_single_provider", side_effect=mock_call_side_effect
):
result = client.chat([{"role": "user", "content": "test"}])
assert result.content == "fallback success"
assert call_count["n"] == 2
def test_fallback_all_fail_raises(self):
"""全部供应商都失败时抛出聚合错误。"""
settings = MockSettings(
llm_provider="dashscope",
llm_api_key="key-main",
llm_model="qwen-plus",
llm_fallback_models="gpt-4o-mini",
llm_fallback_providers="openai",
llm_fallback_api_keys="key-openai",
)
client = LLMClient(settings)
with patch.object(
client, "_call_single_provider", side_effect=LLMError("all fail")
):
with pytest.raises(LLMError, match="全部 2 个"):
client.chat([{"role": "user", "content": "test"}])