Files
baozaotumao2025 14722be770 first commit
2026-07-18 21:00:26 +08:00

214 lines
7.9 KiB
Python

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}")