214 lines
7.9 KiB
Python
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}")
|