Add local model provider boundary
This commit is contained in:
@@ -17,6 +17,7 @@ from fastapi.responses import StreamingResponse
|
||||
from pydantic import BaseModel, Field, model_validator
|
||||
|
||||
from app import database
|
||||
from app.model_provider import get_model_provider
|
||||
from app import workspace
|
||||
|
||||
logger = logging.getLogger("sovereignai.audio")
|
||||
@@ -58,7 +59,7 @@ class AgentRequest(BaseModel):
|
||||
task: str = Field(description="مهمة قصيرة للوكيل المحلي")
|
||||
model: str | None = Field(
|
||||
default=None,
|
||||
description="اسم نموذج Ollama؛ اتركه فارغًا لاستخدام النموذج الافتراضي",
|
||||
description="اسم النموذج المتاح لدى المزوّد المحلي؛ اتركه فارغًا لاستخدام الافتراضي",
|
||||
)
|
||||
|
||||
|
||||
@@ -94,7 +95,7 @@ class ConversationWrite(BaseModel):
|
||||
class WebReadRequest(BaseModel):
|
||||
url: str = Field(min_length=8, max_length=2048, description="رابط صفحة ويب عامة تريد تحليلها")
|
||||
question: str = Field(default="لخّص محتوى الصفحة وأهم نقاطها.", min_length=1, max_length=2000)
|
||||
model: str | None = Field(default=None, description="نموذج Ollama المحلي؛ اتركه فارغًا للنموذج الافتراضي")
|
||||
model: str | None = Field(default=None, description="نموذج المزوّد المحلي؛ اتركه فارغًا للنموذج الافتراضي")
|
||||
|
||||
|
||||
class _PageText(HTMLParser):
|
||||
@@ -223,26 +224,20 @@ def validate_conversation_id(value: str) -> str:
|
||||
|
||||
|
||||
async def get_completion(
|
||||
payload: dict[str, Any], base_url: str, *, timeout_seconds: float = 180.0
|
||||
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"{base_url}/chat/completions", json=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
|
||||
provider = get_model_provider()
|
||||
return await provider.complete(payload, timeout_seconds=timeout_seconds)
|
||||
|
||||
|
||||
@app.get("/health")
|
||||
def health() -> dict[str, Any]:
|
||||
provider = get_model_provider()
|
||||
return {
|
||||
"status": "ok",
|
||||
"model": os.getenv("LOCAL_MODEL", "gemma4:e2b"),
|
||||
"backend": os.getenv("LOCAL_LLM_BASE_URL", "http://127.0.0.1:11434/v1"),
|
||||
"provider": provider.name,
|
||||
"model": provider.default_model,
|
||||
"backend": provider.base_url,
|
||||
"groq_transcription": "configured" if os.getenv("GROQ_API_KEY") else "not_configured",
|
||||
"conversation_database": "sqlite",
|
||||
"workspace_agent": "enabled" if workspace.configured_root() else "not_configured",
|
||||
@@ -251,18 +246,10 @@ def health() -> dict[str, Any]:
|
||||
|
||||
@app.get("/v1/models")
|
||||
async def list_local_models() -> dict[str, Any]:
|
||||
"""List models installed in the configured local Ollama instance."""
|
||||
base_url = os.getenv("LOCAL_LLM_BASE_URL", "http://127.0.0.1:11434/v1").rstrip("/")
|
||||
ollama_base = base_url[:-3] if base_url.endswith("/v1") else 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
|
||||
models = [item["name"] for item in payload.get("models", []) if isinstance(item, dict) and item.get("name")]
|
||||
active = os.getenv("LOCAL_MODEL", "gemma4:e2b")
|
||||
"""List models available from the configured provider."""
|
||||
provider = get_model_provider()
|
||||
models = await provider.list_models()
|
||||
active = provider.default_model
|
||||
if active not in models:
|
||||
models.insert(0, active)
|
||||
return {"data": [{"id": model, "object": "model"} for model in models]}
|
||||
@@ -278,7 +265,7 @@ async def ask_workspace(request: WorkspaceAgentRequest) -> dict[str, Any]:
|
||||
if not files:
|
||||
raise HTTPException(status_code=404, detail="لم أجد نصوصًا مطابقة في ملفات مساحة العمل.")
|
||||
context = "\n\n".join(f"--- ملف: {name} ---\n{content}" for name, content in files)
|
||||
model = request.model or os.getenv("LOCAL_MODEL", "gemma4:e2b")
|
||||
model = request.model or get_model_provider().default_model
|
||||
payload: dict[str, Any] = {
|
||||
"model": model,
|
||||
"messages": [
|
||||
@@ -298,13 +285,7 @@ async def ask_workspace(request: WorkspaceAgentRequest) -> dict[str, Any]:
|
||||
],
|
||||
"stream": False,
|
||||
}
|
||||
if model.lower().startswith("gemma4"):
|
||||
payload["reasoning_effort"] = "none"
|
||||
completion = await get_completion(
|
||||
payload,
|
||||
os.getenv("LOCAL_LLM_BASE_URL", "http://127.0.0.1:11434/v1").rstrip("/"),
|
||||
timeout_seconds=600.0,
|
||||
)
|
||||
completion = await get_completion(payload, timeout_seconds=600.0)
|
||||
return {
|
||||
"task": request.task,
|
||||
"tool": "workspace-search-readonly",
|
||||
@@ -322,7 +303,7 @@ def get_local_user() -> dict[str, str]:
|
||||
|
||||
|
||||
def chat_payload(request: ChatRequest, *, stream: bool) -> dict[str, Any]:
|
||||
model = request.model or os.getenv("LOCAL_MODEL", "gemma4:e2b")
|
||||
model = request.model or get_model_provider().default_model
|
||||
payload: dict[str, Any] = {
|
||||
"model": model,
|
||||
"messages": (
|
||||
@@ -332,7 +313,7 @@ def chat_payload(request: ChatRequest, *, stream: bool) -> dict[str, Any]:
|
||||
{
|
||||
"role": "system",
|
||||
"content": (
|
||||
"أنت مساعد ذكاء اصطناعي محلي يعمل عبر Ollama على جهاز المستخدم. "
|
||||
"أنت مساعد ذكاء اصطناعي يعمل عبر مزوّد النموذج المحلي المضبوط على جهاز المستخدم. "
|
||||
"أجب بالعربية الواضحة وباختصار مناسب. إذا سُئلت أين أنت، أجب بهذه الصياغة: "
|
||||
"أنا مساعد ذكاء اصطناعي يعمل على جهازك، ولا أملك وجودًا جسديًا أو موقع GPS. "
|
||||
"ولا تدّع معرفة "
|
||||
@@ -347,54 +328,27 @@ def chat_payload(request: ChatRequest, *, stream: bool) -> dict[str, Any]:
|
||||
),
|
||||
"stream": stream,
|
||||
}
|
||||
# Disable Gemma 4 thinking so Ollama places the answer in `content` for clients.
|
||||
if model.lower().startswith("gemma4"):
|
||||
payload["reasoning_effort"] = "none"
|
||||
return payload
|
||||
|
||||
|
||||
@app.post("/v1/chat/completions")
|
||||
async def chat(request: ChatRequest) -> dict[str, Any]:
|
||||
"""ترحيل طلب المحادثة إلى النموذج المحلي."""
|
||||
base_url = os.getenv("LOCAL_LLM_BASE_URL", "http://127.0.0.1:11434/v1").rstrip("/")
|
||||
"""إرسال طلب المحادثة إلى مزوّد النموذج المضبوط."""
|
||||
payload = chat_payload(request, stream=False)
|
||||
return await get_completion(payload, base_url)
|
||||
return await get_completion(payload)
|
||||
|
||||
|
||||
@app.post("/v1/chat/stream")
|
||||
async def chat_stream(request: ChatRequest) -> StreamingResponse:
|
||||
"""Pass Ollama token deltas to clients as newline-delimited JSON."""
|
||||
base_url = os.getenv("LOCAL_LLM_BASE_URL", "http://127.0.0.1:11434/v1").rstrip("/")
|
||||
"""Pass provider token deltas to clients as newline-delimited JSON."""
|
||||
provider = get_model_provider()
|
||||
payload = chat_payload(request, stream=True)
|
||||
|
||||
async def events():
|
||||
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"{base_url}/chat/completions", json=payload
|
||||
) as response:
|
||||
if response.status_code >= 400:
|
||||
detail = (await response.aread()).decode("utf-8", "replace")[:400]
|
||||
yield json.dumps({"error": f"Ollama {response.status_code}: {detail}"}) + "\n"
|
||||
return
|
||||
async for line in response.aiter_lines():
|
||||
if not line.startswith("data:"):
|
||||
continue
|
||||
raw = line[5:].strip()
|
||||
if raw == "[DONE]":
|
||||
yield '{"done":true}\n'
|
||||
return
|
||||
try:
|
||||
event = json.loads(raw)
|
||||
delta = event["choices"][0].get("delta", {}).get("content")
|
||||
except (ValueError, KeyError, IndexError, TypeError):
|
||||
continue
|
||||
if delta:
|
||||
yield json.dumps({"delta": delta}, ensure_ascii=False) + "\n"
|
||||
yield '{"done":true}\n'
|
||||
except httpx.RequestError:
|
||||
yield json.dumps({"error": "تعذر الاتصال بـ Ollama المحلي."}, ensure_ascii=False) + "\n"
|
||||
async for event in provider.stream(payload):
|
||||
yield json.dumps(event, ensure_ascii=False) + "\n"
|
||||
if event.get("done") or event.get("error"):
|
||||
return
|
||||
|
||||
return StreamingResponse(
|
||||
events(),
|
||||
@@ -503,8 +457,7 @@ async def run_agent(request: AgentRequest) -> dict[str, Any]:
|
||||
except (ValueError, SyntaxError, ZeroDivisionError):
|
||||
pass
|
||||
|
||||
base_url = os.getenv("LOCAL_LLM_BASE_URL", "http://127.0.0.1:11434/v1").rstrip("/")
|
||||
model = request.model or os.getenv("LOCAL_MODEL", "gemma4:e2b")
|
||||
model = request.model or get_model_provider().default_model
|
||||
payload = {
|
||||
"model": model,
|
||||
"messages": [
|
||||
@@ -513,9 +466,7 @@ async def run_agent(request: AgentRequest) -> dict[str, Any]:
|
||||
],
|
||||
"stream": False,
|
||||
}
|
||||
if model.lower().startswith("gemma4"):
|
||||
payload["reasoning_effort"] = "none"
|
||||
completion = await get_completion(payload, base_url)
|
||||
completion = await get_completion(payload)
|
||||
return {
|
||||
"task": request.task,
|
||||
"tool": "local-llm",
|
||||
@@ -528,7 +479,7 @@ async def run_agent(request: AgentRequest) -> dict[str, Any]:
|
||||
async def read_web_page(request: WebReadRequest) -> dict[str, Any]:
|
||||
"""Fetch a user-provided public web page and ask the local model about its text."""
|
||||
source_url, title, page_text = await _read_public_page(request.url)
|
||||
model = request.model or os.getenv("LOCAL_MODEL", "gemma4:e2b")
|
||||
model = request.model or get_model_provider().default_model
|
||||
payload: dict[str, Any] = {
|
||||
"model": model,
|
||||
"messages": [
|
||||
@@ -553,11 +504,8 @@ async def read_web_page(request: WebReadRequest) -> dict[str, Any]:
|
||||
],
|
||||
"stream": False,
|
||||
}
|
||||
if model.lower().startswith("gemma4"):
|
||||
payload["reasoning_effort"] = "none"
|
||||
completion = await get_completion(
|
||||
payload,
|
||||
os.getenv("LOCAL_LLM_BASE_URL", "http://127.0.0.1:11434/v1").rstrip("/"),
|
||||
timeout_seconds=600.0,
|
||||
)
|
||||
return {
|
||||
|
||||
@@ -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