Artifex/model_router/providers.py
2026-08-15 13:50:24 +07:00

154 lines
5.8 KiB
Python

from __future__ import annotations
import json
import subprocess
import urllib.error
import urllib.request
from dataclasses import dataclass
from typing import Any
from control_plane.resources.models import Resource
from model_router.router import ModelRequestContract, ModelResponseContract
class ProviderError(RuntimeError):
pass
def _extract_json(text: str) -> dict[str, Any]:
try:
value = json.loads(text)
except json.JSONDecodeError as exc:
raise ProviderError("Provider returned malformed JSON") from exc
if not isinstance(value, dict):
raise ProviderError("Provider JSON response must be an object")
return value
def extract_json_object(text: str) -> dict[str, Any]:
try:
return _extract_json(text)
except ProviderError:
start = text.find("{")
end = text.rfind("}")
if start == -1 or end == -1 or end <= start:
raise
return _extract_json(text[start : end + 1])
@dataclass
class SolProvider:
resource: Resource
provider_name: str = "opencode"
def complete(self, request: ModelRequestContract) -> ModelResponseContract:
config = self.resource.config
compute = self.resource.compute
ssh_alias = (compute.config if compute else {}).get("ssh_alias", config.get("ssh_alias", "spark"))
timeout = int(config.get("timeout_seconds", 120))
remote_command = config.get("command", "opencode run --json --no-repo --stdin")
payload = {
"mode": "reasoning_only",
"output_schema": "json_object",
"prompt": request.prompt,
"token_budget": request.token_budget,
}
completed = subprocess.run(
["ssh", str(ssh_alias), str(remote_command)],
input=json.dumps(payload),
capture_output=True,
text=True,
timeout=timeout,
check=False,
)
if completed.returncode != 0:
raise ProviderError(completed.stderr.strip() or "Sol provider failed")
data = _extract_json(completed.stdout)
content = data.get("content") or data.get("response") or data.get("plan")
if isinstance(content, (dict, list)):
content = json.dumps(content)
if not isinstance(content, str):
raise ProviderError("Sol provider response missing string content")
return ModelResponseContract(
model=self.resource.name,
content=content,
metadata={"provider": self.provider_name, "usage": data.get("usage", {})},
)
def health(self) -> str:
compute = self.resource.compute
ssh_alias = (compute.config if compute else {}).get("ssh_alias", self.resource.config.get("ssh_alias", "spark"))
try:
completed = subprocess.run(
["ssh", str(ssh_alias), "true"], capture_output=True, text=True, timeout=10, check=False
)
except Exception:
return "UNAVAILABLE"
return "AVAILABLE" if completed.returncode == 0 else "UNAVAILABLE"
@dataclass
class QwenProvider:
resource: Resource
provider_name: str = "local_inference"
def complete(self, request: ModelRequestContract) -> ModelResponseContract:
config = self.resource.config
url = str(config.get("endpoint_url", "http://localhost:8000/v1/chat/completions"))
timeout = int(config.get("timeout_seconds", 120))
body = {
"model": config.get("model", self.resource.name),
"messages": [{"role": "user", "content": request.prompt}],
"max_tokens": request.token_budget,
"temperature": config.get("temperature", 0),
}
if config.get("response_format"):
body["response_format"] = config["response_format"]
if config.get("extra_body"):
body.update(config["extra_body"])
http_request = urllib.request.Request(
url,
data=json.dumps(body).encode("utf-8"),
headers={"Content-Type": "application/json"},
method="POST",
)
try:
with urllib.request.urlopen(http_request, timeout=timeout) as response:
data = json.loads(response.read().decode("utf-8"))
except (urllib.error.URLError, TimeoutError, json.JSONDecodeError) as exc:
raise ProviderError(f"Qwen provider failed: {exc}") from exc
choices = data.get("choices", [])
content = ""
if choices:
message = choices[0].get("message", {})
content = message.get("content", "")
if not isinstance(content, str) or not content:
raise ProviderError("Qwen provider response missing content")
return ModelResponseContract(
model=str(data.get("model", self.resource.name)),
content=content,
metadata={"provider": self.provider_name, "usage": data.get("usage", {})},
)
def health(self) -> str:
base_url = str(self.resource.config.get("health_url", self.resource.config.get("endpoint_url", ""))).replace(
"/v1/chat/completions", "/health"
)
if not base_url:
return "UNAVAILABLE"
try:
with urllib.request.urlopen(base_url, timeout=5) as response:
return "AVAILABLE" if 200 <= response.status < 500 else "DEGRADED"
except Exception:
return "UNAVAILABLE"
def providers_from_resources() -> dict[str, object]:
providers: dict[str, object] = {}
sol = Resource.objects.filter(is_active=True, provider="opencode").first()
qwen = Resource.objects.filter(is_active=True, provider="local_inference").first()
if sol is not None:
providers["sol"] = SolProvider(sol)
if qwen is not None:
providers["qwen"] = QwenProvider(qwen)
return providers