first commit
This commit is contained in:
@@ -0,0 +1,213 @@
|
||||
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}")
|
||||
Reference in New Issue
Block a user