"""Provider boundary for local and future model backends.""" from __future__ import annotations import asyncio import json import os from collections.abc import AsyncIterator from typing import Any, Protocol import httpx from fastapi import HTTPException class ModelProvider(Protocol): name: str default_model: str base_url: str def prepare_payload(self, payload: dict[str, Any]) -> dict[str, Any]: ... async def complete( self, payload: dict[str, Any], *, timeout_seconds: float = 180.0 ) -> dict[str, Any]: ... async def stream( self, payload: dict[str, Any] ) -> AsyncIterator[dict[str, Any]]: ... async def list_models(self) -> list[str]: ... async def describe_models( self, model_names: list[str] ) -> dict[str, dict[str, Any]]: ... class OllamaProvider: name = "ollama" def __init__(self, base_url: str, default_model: str) -> None: self.base_url = base_url.rstrip("/") self.default_model = default_model def prepare_payload(self, payload: dict[str, Any]) -> dict[str, Any]: prepared = dict(payload) model = str(prepared.get("model", self.default_model)) prepared.setdefault("model", model) # Gemma 4 via Ollama otherwise returns its reasoning in a separate field. if model.lower().startswith("gemma4"): prepared.setdefault("reasoning_effort", "none") return prepared async def complete( self, payload: dict[str, Any], *, timeout_seconds: float = 180.0 ) -> dict[str, Any]: try: async with httpx.AsyncClient(timeout=timeout_seconds) as client: response = await client.post( f"{self.base_url}/chat/completions", json=self.prepare_payload(payload), ) response.raise_for_status() return response.json() except httpx.HTTPStatusError as exc: detail = exc.response.text[:400] or "رفض خادم النموذج الطلب." raise HTTPException( status_code=502, detail=f"خطأ من خادم النموذج المحلي: {detail}", ) from exc except httpx.TimeoutException as exc: raise HTTPException( status_code=504, detail="انتهت مهلة انتظار خادم النموذج المحلي.", ) from exc except httpx.RequestError as exc: raise HTTPException( status_code=503, detail="تعذر الاتصال بـ Ollama المحلي على العنوان المضبوط.", ) from exc except ValueError as exc: raise HTTPException( status_code=502, detail="أعاد Ollama استجابة JSON غير صالحة." ) from exc async def stream( self, payload: dict[str, Any] ) -> AsyncIterator[dict[str, Any]]: try: timeout = httpx.Timeout(connect=15.0, read=None, write=30.0, pool=30.0) async with httpx.AsyncClient(timeout=timeout) as client: async with client.stream( "POST", f"{self.base_url}/chat/completions", json=self.prepare_payload(payload), ) as response: if response.status_code >= 400: detail = (await response.aread()).decode("utf-8", "replace")[:400] yield {"error": f"Ollama {response.status_code}: {detail}"} return async for line in response.aiter_lines(): if not line.startswith("data:"): continue raw = line[5:].strip() if raw == "[DONE]": yield {"done": True} return try: event = json.loads(raw) delta = event["choices"][0].get("delta", {}).get("content") except (ValueError, KeyError, IndexError, TypeError): continue if delta: yield {"delta": delta} yield {"done": True} except httpx.RequestError: yield {"error": "تعذر الاتصال بـ Ollama المحلي."} async def list_models(self) -> list[str]: ollama_base = ( self.base_url[:-3] if self.base_url.endswith("/v1") else self.base_url ) try: async with httpx.AsyncClient(timeout=10.0) as client: response = await client.get(f"{ollama_base}/api/tags") response.raise_for_status() payload = response.json() except (httpx.HTTPError, ValueError) as exc: raise HTTPException( status_code=503, detail="تعذر جلب قائمة النماذج من Ollama المحلي.", ) from exc return [ item["name"] for item in payload.get("models", []) if isinstance(item, dict) and item.get("name") ] async def describe_models( self, model_names: list[str] ) -> dict[str, dict[str, Any]]: """Read capability metadata advertised by the installed Ollama models.""" ollama_base = ( self.base_url[:-3] if self.base_url.endswith("/v1") else self.base_url ) capability_order = [ "completion", "vision", "audio", "tools", "thinking", "embedding", ] semaphore = asyncio.Semaphore(4) async def describe(client: httpx.AsyncClient, model_name: str): async with semaphore: try: response = await client.post( f"{ollama_base}/api/show", json={"model": model_name, "verbose": False}, ) response.raise_for_status() payload = response.json() raw = ( payload.get("capabilities") if isinstance(payload, dict) else None ) verified = isinstance(raw, list) return model_name, { "verified": verified, "capabilities": [ capability for capability in capability_order if capability in raw ] if verified else [], } except (httpx.HTTPError, ValueError, TypeError): return model_name, {"verified": False, "capabilities": []} try: async with httpx.AsyncClient(timeout=10.0) as client: described = await asyncio.gather( *(describe(client, name) for name in model_names) ) return dict(described) except httpx.HTTPError: return { name: {"verified": False, "capabilities": []} for name in model_names } def get_model_provider() -> ModelProvider: provider_name = os.getenv("MODEL_PROVIDER", "ollama").strip().lower() if provider_name != "ollama": raise HTTPException( status_code=503, detail=f"مزوّد النموذج '{provider_name}' غير مدعوم حاليًا.", ) return OllamaProvider( base_url=os.getenv("LOCAL_LLM_BASE_URL", "http://127.0.0.1:11434/v1"), default_model=os.getenv("LOCAL_MODEL", "gemma4:e2b"), )