"""Provider boundary for local and future model backends.""" from __future__ import annotations 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]: ... 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.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") ] 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"), )