154 lines
5.8 KiB
Python
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
|