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 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: 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)