162 lines
5.8 KiB
Python
162 lines
5.8 KiB
Python
import os
|
|
from collections.abc import Mapping
|
|
from dataclasses import dataclass
|
|
from typing import Any, Protocol, cast
|
|
|
|
import httpx
|
|
from sqlalchemy import select
|
|
|
|
from app.core.errors import RecoverableAgentError
|
|
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: ...
|
|
|
|
|
|
class EndpointSettings(Protocol):
|
|
endpoint_code: str
|
|
base_url: str
|
|
model_name: str
|
|
secret_ref: str
|
|
|
|
|
|
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
|
|
|
|
|
|
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 RecoverableAgentError("模型响应缺少文本")
|
|
return content
|
|
except httpx.HTTPError as exc:
|
|
raise RecoverableAgentError("模型端点调用失败") from exc
|
|
finally:
|
|
if self._owns_client:
|
|
await client.aclose()
|
|
|
|
|
|
class DatabaseModelGateway:
|
|
"""从当前数据库端点配置动态构造已批准的 OpenAI-compatible Adapter。"""
|
|
|
|
async def generate(self, *, endpoint_code: str, prompt: str, timeout_ms: int) -> str:
|
|
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("模型端点未注册或未激活")
|
|
adapter = OpenAICompatibleGateway({
|
|
endpoint.endpoint_code: cast(EndpointSettings, endpoint)
|
|
})
|
|
return await adapter.generate(
|
|
endpoint_code=endpoint.endpoint_code, prompt=prompt, timeout_ms=timeout_ms
|
|
)
|
|
|
|
|
|
class DatabaseModelEndpointResolver:
|
|
"""为统一意图分类提供当前已激活模型端点快照。"""
|
|
|
|
async def resolve(self, *, agent_type: str, task_type: str) -> list[ModelEndpointConfig]:
|
|
del agent_type, task_type # 路由筛选由发布配置和 ModelRouterService 扩展。
|
|
async with SessionFactory() as session:
|
|
return list(await session.scalars(select(ModelEndpointConfig).where(
|
|
ModelEndpointConfig.status == "active"
|
|
)))
|
|
|
|
|
|
@dataclass(frozen=True)
|
|
class ModelExecution:
|
|
endpoint_code: str
|
|
text: str
|
|
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
|
|
|
|
|
|
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)
|