Add local model provider boundary
This commit is contained in:
@@ -0,0 +1,139 @@
|
||||
"""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"),
|
||||
)
|
||||
Reference in New Issue
Block a user