140 lines
5.3 KiB
Python
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"),
|
|
)
|