import asyncio from abc import ABC, abstractmethod from dataclasses import dataclass from typing import Any import httpx from app.core.config import Settings from app.core.exceptions import ConfigurationError, ProviderError from app.db.repositories.repositories import ProviderConfigRepository from sqlalchemy.orm import Session @dataclass class ChatResponse: content: str tool_calls: list[dict[str, Any]] raw: dict[str, Any] class ModelProvider(ABC): @abstractmethod async def list_models(self) -> list[str]: raise NotImplementedError @abstractmethod async def chat(self, messages: list[dict[str, Any]], tools=None) -> ChatResponse: raise NotImplementedError @property @abstractmethod def endpoint(self) -> str: raise NotImplementedError @property @abstractmethod def model(self) -> str: raise NotImplementedError class OpenAICompatibleProvider(ModelProvider): def __init__( self, *, base_url: str, chat_path: str, models_path: str, model: str, auth_type: str, api_key: str, auth_header: str, auth_prefix: str, timeout: int, retry_count: int, retry_delay: int, verify_ssl: bool, retry_max_delay: int = 60, ): self.base_url = base_url.rstrip("/") self.chat_path = "/" + chat_path.lstrip("/") self.models_path = "/" + models_path.lstrip("/") self._model = model self.timeout = timeout self.retry_count = retry_count self.retry_delay = retry_delay self.retry_max_delay = retry_max_delay self.verify_ssl = verify_ssl self.headers = {"Content-Type": "application/json"} if auth_type != "none": if not api_key: raise ConfigurationError("认证已启用,但 API Key 为空") self.headers[auth_header] = f"{auth_prefix.strip()} {api_key}".strip() @property def endpoint(self) -> str: return self.base_url + self.chat_path @property def model(self) -> str: return self._model async def list_models(self) -> list[str]: response = await self._request("GET", self.base_url + self.models_path) if response.status_code >= 400: raise ProviderError( f"模型列表接口返回 HTTP {response.status_code}", details={"body": response.text[:1000]}, ) data = response.json() items = data.get("data") or data.get("models") or [] result = [] for item in items: if isinstance(item, str): result.append(item) elif isinstance(item, dict): value = item.get("id") or item.get("name") or item.get("model") if value: result.append(str(value)) return result async def chat(self, messages: list[dict[str, Any]], tools=None) -> ChatResponse: if not self._model: raise ConfigurationError("模型名称为空") payload: dict[str, Any] = { "model": self._model, "messages": messages, "temperature": 0, } if tools: payload["tools"] = tools payload["tool_choice"] = "auto" response = await self._request("POST", self.endpoint, json=payload) if response.status_code >= 400: raise ProviderError( f"模型接口返回 HTTP {response.status_code}", details={"body": response.text[:2000]}, ) data = response.json() try: message = data["choices"][0]["message"] except Exception as exc: raise ProviderError("无法解析模型响应", details={"response": data}) from exc return ChatResponse( content=message.get("content") or "", tool_calls=message.get("tool_calls") or [], raw=data, ) async def _request(self, method: str, url: str, **kwargs) -> httpx.Response: async with httpx.AsyncClient(timeout=self.timeout, verify=self.verify_ssl) as client: for attempt in range(self.retry_count + 1): try: response = await client.request(method, url, headers=self.headers, **kwargs) if response.status_code not in {408, 429} and response.status_code < 500: return response if attempt == self.retry_count: return response retry_after = response.headers.get("Retry-After") try: delay = float(retry_after) if retry_after else self.retry_delay * 2**attempt except ValueError: delay = self.retry_delay * 2**attempt delay = min(delay, self.retry_max_delay) except httpx.RequestError as exc: if attempt == self.retry_count: detail = str(exc) or "无详细信息" raise ProviderError( f"模型请求失败: {type(exc).__name__}: {detail}" ) from exc delay = min(self.retry_delay * 2**attempt, self.retry_max_delay) await asyncio.sleep(delay) raise RuntimeError("不可达") class ProviderFactory: def __init__(self, settings: Settings, db: Session | None = None): self.settings = settings self.db = db def create(self, provider: str) -> ModelProvider: profile = ProviderConfigRepository(self.db).get_by_type(provider) if self.db else None if profile: return OpenAICompatibleProvider( base_url=profile.endpoint, chat_path=profile.chat_path, models_path=profile.models_path, model=profile.model_name, auth_type=profile.auth_type, api_key=profile.api_key, auth_header=profile.auth_header, auth_prefix=profile.auth_prefix, timeout=self.settings.request_timeout_seconds, retry_count=self.settings.request_retry_count, retry_delay=self.settings.request_retry_delay_seconds, retry_max_delay=self.settings.request_retry_max_delay_seconds, verify_ssl=profile.verify_ssl, ) if provider == "target": s = self.settings return OpenAICompatibleProvider( base_url=s.target_base_url, chat_path=s.target_chat_path, models_path=s.target_models_path, model=s.target_model, auth_type=s.target_auth_type, api_key=s.target_api_key, auth_header=s.target_auth_header, auth_prefix=s.target_auth_prefix, timeout=s.request_timeout_seconds, retry_count=s.request_retry_count, retry_delay=s.request_retry_delay_seconds, retry_max_delay=s.request_retry_max_delay_seconds, verify_ssl=s.request_verify_ssl, ) if provider == "judge": s = self.settings return OpenAICompatibleProvider( base_url=s.judge_base_url, chat_path=s.judge_chat_path, models_path=s.judge_models_path, model=s.judge_model, auth_type=s.judge_auth_type, api_key=s.judge_api_key, auth_header=s.judge_auth_header, auth_prefix=s.judge_auth_prefix, timeout=s.request_timeout_seconds, retry_count=s.request_retry_count, retry_delay=s.request_retry_delay_seconds, retry_max_delay=s.request_retry_max_delay_seconds, verify_ssl=s.request_verify_ssl, ) raise ConfigurationError(f"未知 provider={provider}")