feat(llm): 多模型 fallback 支持
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:
@@ -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
|
||||||
|
|||||||
@@ -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=
|
||||||
|
|||||||
@@ -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)。
|
||||||
|
|||||||
@@ -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,
|
||||||
|
|||||||
@@ -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"}])
|
||||||
|
|||||||
Reference in New Issue
Block a user