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