Files
group_fqcd_jr/app/service/model_gateway.py
T
2026-09-11 15:38:11 +08:00

334 lines
14 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
import logging
import os
from collections.abc import Mapping
from dataclasses import dataclass
from typing import Any, Protocol
import httpx
from sqlalchemy import select
from app.core.errors import (
DependencyUnavailableError,
RecoverableAgentError,
UpstreamTimeoutError,
)
from app.infrastructure.db import SessionFactory
from app.model.configuration import ModelEndpointConfig
class ModelGateway(Protocol):
async def generate(self, *, endpoint_code: str, prompt: str, timeout_ms: int) -> str: ...
async def embed(self, *, endpoint_code: str, text: str, timeout_ms: int) -> list[float]: ...
class EndpointSettings(Protocol):
endpoint_code: str
base_url: str
model_name: str
secret_ref: str
logger = logging.getLogger(__name__)
class EnvironmentSecretResolver:
def resolve(self, secret_ref: str) -> str:
if not secret_ref.startswith("env:"):
raise RecoverableAgentError("模型密钥必须使用 env: 引用")
name = secret_ref.removeprefix("env:")
value = os.getenv(name)
if not value:
raise RecoverableAgentError("模型密钥未配置")
return value
#: `task_type` → 端点必须具备的能力标签。
#:
#: 取值来自全仓实际调用点与**现库端点实际声明的能力**(实测:
#: `embedding-primary` = `["embedding"]`、`chat-primary` = `["chat"]`):
#: `embedding`(向量化:`bootstrap.get_memory_embedding_service`、`knowledge_vector_worker`)、
#: `chat`(生成/复述合并:`customer_service`)、`intent_classification`(`BaseAgent.classify_intent`)、
#: `memory_extraction` / `text_generation`(结构化文本生成)。
#:
#: ⚠️ `intent_classification` 必须映射到 `chat`,**不能**映射成同名的
#: `intent_classification`:现库没有任何端点声明该能力,同名映射会筛出空集,
#: 依赖"回退到全部端点"才不至于失败——那是把配置缺陷掩盖成巧合。
#: 未列入的 `task_type` 不做过滤(无法判断该要哪种能力),但会**告警留痕**(见 `resolve`)。
REQUIRED_CAPABILITY_BY_TASK_TYPE: dict[str, str] = {
"embedding": "embedding",
"chat": "chat",
"intent_classification": "chat",
"memory_extraction": "text_generation",
# 风控的几处 task_type:它们要的都是"能生成文本"的端点,与 memory_extraction 同理。
# 不登记就会落到"未映射 → 返回全部端点"的分支,而能否选到文本端点就取决于
# `model_endpoint_config` 的**行顺序**——实测风控能跑通,仅仅因为 deepseek-flash(id=3)
# 恰好排在 qwen-embedding(id=5) 前面。这种"靠数据顺序才对"的隐式依赖必须消掉。
"risk_agent_chat": "text_generation",
"risk_analysis": "text_generation",
"risk_script": "text_generation",
"risk_summary": "text_generation",
"daily_report_suggestion": "text_generation",
"text_generation": "text_generation",
}
class OpenAICompatibleGateway:
"""供应商无关的 Chat Completions Adapter;密钥只从 secret_ref 解析。"""
def __init__(
self,
endpoints: Mapping[str, EndpointSettings],
*,
secret_resolver: EnvironmentSecretResolver | None = None,
client: httpx.AsyncClient | None = None,
) -> None:
self.endpoints = endpoints
self.secret_resolver = secret_resolver or EnvironmentSecretResolver()
self.client = client
self._owns_client = client is None
async def generate(self, *, endpoint_code: str, prompt: str, timeout_ms: int) -> str:
endpoint = self.endpoints.get(endpoint_code)
if endpoint is None:
raise RecoverableAgentError("模型端点未注册")
token = self.secret_resolver.resolve(endpoint.secret_ref)
url = endpoint.base_url.rstrip("/") + "/chat/completions"
payload = {
"model": endpoint.model_name,
"messages": [{"role": "user", "content": prompt}],
"temperature": 0,
}
client = self.client or httpx.AsyncClient()
try:
response = await client.post(
url,
headers={"Authorization": f"Bearer {token}", "Content-Type": "application/json"},
json=payload,
timeout=httpx.Timeout(timeout_ms / 1000),
)
response.raise_for_status()
body: Any = response.json()
content = body.get("choices", [{}])[0].get("message", {}).get("content")
if not isinstance(content, str) or not content.strip():
raise DependencyUnavailableError("模型响应缺少文本")
return content
except httpx.TimeoutException as exc:
raise UpstreamTimeoutError("模型端点调用超时") from exc
except httpx.HTTPError as exc:
raise DependencyUnavailableError("模型端点调用失败") from exc
finally:
if self._owns_client:
await client.aclose()
async def embed(self, *, endpoint_code: str, text: str, timeout_ms: int) -> list[float]:
"""OpenAI-compatible Embeddings Adapter;密钥同样只从 secret_ref 解析。"""
endpoint = self.endpoints.get(endpoint_code)
if endpoint is None:
raise RecoverableAgentError("模型端点未注册")
token = self.secret_resolver.resolve(endpoint.secret_ref)
url = endpoint.base_url.rstrip("/") + "/embeddings"
payload = {"model": endpoint.model_name, "input": text}
client = self.client or httpx.AsyncClient()
try:
response = await client.post(
url,
headers={"Authorization": f"Bearer {token}", "Content-Type": "application/json"},
json=payload,
timeout=httpx.Timeout(timeout_ms / 1000),
)
response.raise_for_status()
body: Any = response.json()
vector = _first_embedding(body)
if vector is None:
raise DependencyUnavailableError("向量响应缺少 embedding")
return vector
except httpx.TimeoutException as exc:
raise UpstreamTimeoutError("向量端点调用超时") from exc
except httpx.HTTPError as exc:
raise DependencyUnavailableError("向量端点调用失败") from exc
finally:
if self._owns_client:
await client.aclose()
def _first_embedding(body: Any) -> list[float] | None:
"""取 OpenAI-compatible 响应的首条向量;结构不符时返回 None,由调用方失败关闭。"""
data = body.get("data") if isinstance(body, dict) else None
if not isinstance(data, list) or not data:
return None
entry = data[0]
raw = entry.get("embedding") if isinstance(entry, dict) else None
if not isinstance(raw, list) or not raw:
return None
try:
return [float(value) for value in raw]
except (TypeError, ValueError):
return None
class DatabaseModelGateway:
"""从当前数据库端点配置动态构造已批准的 OpenAI-compatible Adapter。"""
async def _endpoint(self, endpoint_code: str) -> ModelEndpointConfig:
async with SessionFactory() as session:
endpoint = await session.scalar(select(ModelEndpointConfig).where(
ModelEndpointConfig.endpoint_code == endpoint_code,
ModelEndpointConfig.status == "active",
))
if endpoint is None:
raise RecoverableAgentError("模型端点未注册或未激活")
return endpoint
async def generate(self, *, endpoint_code: str, prompt: str, timeout_ms: int) -> str:
endpoint = await self._endpoint(endpoint_code)
adapter = OpenAICompatibleGateway({endpoint.endpoint_code: endpoint})
return await adapter.generate(
endpoint_code=endpoint.endpoint_code, prompt=prompt, timeout_ms=timeout_ms
)
async def embed(self, *, endpoint_code: str, text: str, timeout_ms: int) -> list[float]:
endpoint = await self._endpoint(endpoint_code)
adapter = OpenAICompatibleGateway({endpoint.endpoint_code: endpoint})
return await adapter.embed(
endpoint_code=endpoint.endpoint_code, text=text, timeout_ms=timeout_ms
)
#: 历史别名,指向上面的唯一映射表。保留名字是为了不破坏既有引用;
#: **两个名字不要各存一份内容**——那正是"修复被后续合并悄悄回退"的成因。
TASK_CAPABILITY = REQUIRED_CAPABILITY_BY_TASK_TYPE
class DatabaseModelEndpointResolver:
"""按 `task_type` 筛选当前已激活的模型端点。
修复的缺陷:原实现 `del agent_type, task_type` 后返回**全部** active 端点,而
`ModelDispatchService` 的 `generate`/`embed` 走的是**不同上游路径**
(`/chat/completions` 与 `/embeddings`),且只按顺序尝试前 `max(1, max_attempts)`
(默认 2)个。两者叠加的后果是「能不能选到支持该任务的端点」取决于端点在表里的顺序——
实测表现为每次 embedding 都先拿文本生成端点失败一次再落到真正的向量端点,
端点一多就会**耗尽尝试次数**直接失败。
筛选策略保持保守:只在确实筛到端点时收窄;未映射的 task_type、以及「一个都没声明
该能力」的配置缺口,都退回全部端点,让故障表现为**调用失败**而不是解析为空。
"""
async def resolve(self, *, agent_type: str, task_type: str) -> list[ModelEndpointConfig]:
del agent_type # 按 agent_type 分配端点的路由仍由发布配置与 ModelRouterService 扩展。
required = REQUIRED_CAPABILITY_BY_TASK_TYPE.get(task_type)
async with SessionFactory() as session:
endpoints = list(await session.scalars(select(ModelEndpointConfig).where(
ModelEndpointConfig.status == "active"
)))
if required is None:
# 未登记的 task_type 仍退回全部端点(保守策略:让故障表现为调用失败、
# 而不是解析为空),但必须留下痕迹。静默退回会让"选端点靠表行顺序"这类问题
# 在下游以"偶发调用失败"的形式冒出来,极难定位。
logger.warning(
"模型端点筛选:task_type=%r 未登记能力映射,退回全部 active 端点;"
"请在 REQUIRED_CAPABILITY_BY_TASK_TYPE 中补上它对应的能力",
task_type,
)
return endpoints
return endpoints
matched = [
endpoint for endpoint in endpoints
if isinstance(endpoint.capabilities, list) and required in endpoint.capabilities
]
# `capabilities` 为 NULL/空 的端点绝不会被 `matched` 选中(不能裸奔到错误的网关方法上);
# 但整批都没声明该能力时退回全部端点,避免把这个配置缺口伪装成"没有可用端点"。
return matched or endpoints
@dataclass(frozen=True)
class ModelExecution:
endpoint_code: str
text: str
attempts: int
degraded: bool = False
@dataclass(frozen=True)
class ModelEmbedding:
endpoint_code: str
vector: list[float]
attempts: int
degraded: bool = False
class ModelDispatchService:
"""Executes only router-approved endpoints and falls back in declared order."""
def __init__(self, gateway: ModelGateway) -> None:
self.gateway = gateway
async def generate(
self,
endpoints: list[Any],
prompt: str,
*,
max_attempts: int = 2,
) -> ModelExecution:
last_error: Exception | None = None
attempts = 0
for endpoint in endpoints[: max(1, max_attempts)]:
attempts += 1
try:
text = await self.gateway.generate(
endpoint_code=endpoint.endpoint_code,
prompt=prompt,
timeout_ms=endpoint.timeout_ms,
)
return ModelExecution(endpoint.endpoint_code, text, attempts, attempts > 1)
except Exception as exc:
last_error = exc
raise RecoverableAgentError("所有已批准模型端点调用失败") from last_error
async def embed(
self, endpoints: list[Any], text: str, *, max_attempts: int = 2
) -> ModelEmbedding:
"""与文本生成同族的受控降级:按声明顺序尝试,成功即返回并标记是否降级。"""
last_error: Exception | None = None
attempts = 0
for endpoint in endpoints[: max(1, max_attempts)]:
attempts += 1
try:
vector = await self.gateway.embed(
endpoint_code=endpoint.endpoint_code,
text=text,
timeout_ms=endpoint.timeout_ms,
)
return ModelEmbedding(endpoint.endpoint_code, vector, attempts, attempts > 1)
except Exception as exc:
last_error = exc
raise RecoverableAgentError("所有已批准向量端点调用失败") from last_error
class ModelGenerationService:
"""业务 Agent 的唯一模型生成入口。路由结果必须先由 ModelRouterService 给出。"""
def __init__(self, dispatch: ModelDispatchService) -> None:
self.dispatch = dispatch
async def generate(
self, endpoints: list[Any], prompt: str, *, max_attempts: int = 2
) -> ModelExecution:
if not endpoints:
raise RecoverableAgentError("没有可用的已批准模型端点")
return await self.dispatch.generate(endpoints, prompt, max_attempts=max_attempts)
class ModelEmbeddingService:
"""文本向量化入口(记忆语义召回等只读用途),与生成同样必须先经过端点解析。"""
def __init__(self, dispatch: ModelDispatchService) -> None:
self.dispatch = dispatch
async def embed(
self, endpoints: list[Any], text: str, *, max_attempts: int = 2
) -> ModelEmbedding:
if not endpoints:
raise RecoverableAgentError("没有可用的已批准向量端点")
return await self.dispatch.embed(endpoints, text, max_attempts=max_attempts)