Files
PromptCR-Lab/backend/app/model_adapters/base.py
T
2026-09-19 12:54:45 +08:00

94 lines
3.1 KiB
Python

"""Abstract base class for model adapters."""
from abc import ABC, abstractmethod
import asyncio
import time
from dataclasses import dataclass
from typing import Any, Dict, Optional
import httpx
@dataclass
class ChatResponse:
text: str
token_usage: Dict[str, int]
latency_ms: float
class ModelAdapter(ABC):
"""Unified interface for LLM vendors.
Subclasses only need to provide base_url, api_key, model name and any
vendor-specific headers. Concurrency and retry logic are inherited.
"""
def __init__(
self,
api_key: str,
model: str,
base_url: str,
concurrency: int = 5,
max_retries: int = 3,
timeout: float = 120.0,
):
self.api_key = api_key
self.model = model
self.base_url = base_url.rstrip("/")
self.semaphore = asyncio.Semaphore(concurrency)
self.max_retries = max_retries
self.timeout = timeout
@abstractmethod
def _build_payload(self, prompt: str, params: Dict[str, Any]) -> Dict[str, Any]:
...
@abstractmethod
def _extract_text(self, data: Dict[str, Any]) -> str:
...
def _extract_token_usage(self, data: Dict[str, Any]) -> Dict[str, int]:
usage = data.get("usage", {})
return {
"prompt_tokens": usage.get("prompt_tokens", 0),
"completion_tokens": usage.get("completion_tokens", 0),
"total_tokens": usage.get("total_tokens", 0),
}
def _headers(self) -> Dict[str, str]:
return {
"Authorization": f"Bearer {self.api_key}",
"Content-Type": "application/json",
}
async def chat(self, prompt: str, params: Optional[Dict[str, Any]] = None) -> ChatResponse:
params = params or {}
payload = self._build_payload(prompt, params)
async with self.semaphore:
last_exception: Optional[Exception] = None
for attempt in range(self.max_retries + 1):
start = time.perf_counter()
try:
async with httpx.AsyncClient(timeout=self.timeout) as client:
response = await client.post(
f"{self.base_url}/chat/completions",
headers=self._headers(),
json=payload,
)
response.raise_for_status()
data = response.json()
latency_ms = (time.perf_counter() - start) * 1000
return ChatResponse(
text=self._extract_text(data),
token_usage=self._extract_token_usage(data),
latency_ms=latency_ms,
)
except Exception as e:
last_exception = e
if attempt < self.max_retries:
wait = 2**attempt
await asyncio.sleep(wait)
raise RuntimeError(
f"Model {self.model} failed after {self.max_retries} retries: {last_exception}"
)