Files
sovereign_ai/SovereignAI-Starter/app/model_provider.py
T

140 lines
5.3 KiB
Python

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