Artifex/model_router/providers.py

298 lines
12 KiB
Python
Raw Permalink Normal View History

2026-08-15 13:50:24 +07:00
from __future__ import annotations
import json
import os
2026-08-15 14:16:27 +07:00
import re
import shlex
2026-08-15 13:50:24 +07:00
import subprocess
2026-08-16 15:52:46 +07:00
import time
2026-08-15 13:50:24 +07:00
import urllib.error
import urllib.request
from collections.abc import Iterator
2026-08-15 13:50:24 +07:00
from dataclasses import dataclass
from typing import Any
from control_plane.resources.models import Resource
from model_router.router import ModelChunk, ModelRequestContract, ModelResponseContract
2026-08-15 13:50:24 +07:00
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])
2026-08-15 14:16:27 +07:00
def _clean_opencode_output(text: str) -> str:
text = re.sub(r"\x1b\[[0-9;?]*[ -/]*[@-~]", "", text)
lines = []
for line in text.splitlines():
stripped = line.strip()
if not stripped or stripped.startswith("> build"):
continue
lines.append(line)
return "\n".join(lines).strip()
2026-08-15 13:50:24 +07:00
@dataclass
class SolProvider:
resource: Resource
provider_name: str = "opencode"
def complete(self, request: ModelRequestContract) -> ModelResponseContract:
config = self.resource.config
compute = self.resource.compute
timeout = int(config.get("timeout_seconds", 120))
2026-08-15 14:16:27 +07:00
command = str(config.get("command", "opencode run"))
transport = str(config.get("transport", "ssh"))
use_pty = bool(config.get("use_pty", False))
working_directory = str(config.get("working_directory") or "").strip() or None
2026-08-15 14:16:27 +07:00
if transport == "local":
argv = shlex.split(command, posix=os.name != "nt")
2026-08-15 14:16:27 +07:00
if use_pty:
shell_command = " ".join(shlex.quote(part) for part in argv)
argv = ["script", "-q", "-e", "-c", shell_command, "/dev/null"]
completed = subprocess.run(
argv,
capture_output=True,
text=True,
timeout=timeout,
check=False,
cwd=working_directory,
input=request.prompt,
)
2026-08-15 14:16:27 +07:00
else:
ssh_alias = (compute.config if compute else {}).get("ssh_alias", config.get("ssh_alias", "spark"))
remote_command = " ".join(shlex.quote(part) for part in shlex.split(command))
if working_directory:
remote_command = f"cd {shlex.quote(working_directory)} && {remote_command}"
2026-08-15 14:16:27 +07:00
completed = subprocess.run(
["ssh", "-T", str(ssh_alias), remote_command],
2026-08-15 14:16:27 +07:00
capture_output=True,
text=True,
timeout=timeout,
check=False,
input=request.prompt,
2026-08-15 14:16:27 +07:00
)
2026-08-15 13:50:24 +07:00
if completed.returncode != 0:
raise ProviderError(completed.stderr.strip() or "Sol provider failed")
2026-08-15 14:16:27 +07:00
content = _clean_opencode_output(completed.stdout)
if not content:
raise ProviderError("Sol provider response missing content")
2026-08-15 13:50:24 +07:00
return ModelResponseContract(
model=self.resource.name,
content=content,
2026-08-15 14:16:27 +07:00
metadata={"provider": self.provider_name, "usage": {}},
2026-08-15 13:50:24 +07:00
)
def stream(self, request: ModelRequestContract) -> Iterator[ModelChunk]:
config = self.resource.config
command = str(config.get("command", "opencode run"))
transport = str(config.get("transport", "ssh"))
working_directory = str(config.get("working_directory") or "").strip() or None
if transport == "local":
argv = shlex.split(command, posix=os.name != "nt")
else:
compute = self.resource.compute
ssh_alias = (compute.config if compute else {}).get(
"ssh_alias", config.get("ssh_alias", "spark")
)
remote_command = " ".join(shlex.quote(part) for part in shlex.split(command))
if working_directory:
remote_command = f"cd {shlex.quote(working_directory)} && {remote_command}"
argv = ["ssh", "-T", str(ssh_alias), remote_command]
process = subprocess.Popen(
argv,
stdout=subprocess.PIPE,
stderr=subprocess.PIPE,
text=True,
bufsize=1,
cwd=working_directory if transport == "local" else None,
stdin=subprocess.PIPE,
)
assert process.stdin is not None
process.stdin.write(request.prompt)
process.stdin.close()
emitted = False
assert process.stdout is not None
for line in process.stdout:
cleaned = _clean_opencode_output(line)
if not cleaned:
continue
emitted = True
yield ModelChunk(cleaned + "\n", {"provider": self.provider_name})
stderr = process.stderr.read().strip() if process.stderr is not None else ""
return_code = process.wait()
if return_code != 0:
raise ProviderError(stderr or "Sol streaming provider failed")
if not emitted:
raise ProviderError("Sol streaming provider response missing content")
2026-08-15 13:50:24 +07:00
def health(self) -> str:
2026-08-15 14:16:27 +07:00
config = self.resource.config
2026-08-15 13:50:24 +07:00
try:
2026-08-15 14:16:27 +07:00
if config.get("transport") == "local":
command = shlex.split(str(config.get("command", "opencode")), posix=os.name != "nt")[0]
2026-08-15 14:16:27 +07:00
completed = subprocess.run([command, "--version"], capture_output=True, text=True, timeout=10, check=False)
else:
compute = self.resource.compute
ssh_alias = (compute.config if compute else {}).get("ssh_alias", self.resource.config.get("ssh_alias", "spark"))
completed = subprocess.run(["ssh", str(ssh_alias), "true"], capture_output=True, text=True, timeout=10, check=False)
2026-08-15 13:50:24 +07:00
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))
2026-08-16 15:52:46 +07:00
attempts = max(1, int(config.get("retry_attempts", 1)))
backoff = float(config.get("retry_backoff_seconds", 1.0))
2026-08-15 13:50:24 +07:00
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"])
2026-08-16 15:52:46 +07:00
failures = []
for attempt in range(1, attempts + 1):
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"))
break
except (urllib.error.HTTPError, urllib.error.URLError, TimeoutError, json.JSONDecodeError) as exc:
failures.append(f"attempt {attempt}/{attempts}: {exc}")
if attempt == attempts:
raise ProviderError("Qwen provider failed after retries: " + " | ".join(failures)) from exc
time.sleep(backoff * attempt)
2026-08-15 13:50:24 +07:00
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,
2026-08-16 15:52:46 +07:00
metadata={
"provider": self.provider_name,
"usage": data.get("usage", {}),
"attempts": len(failures) + 1,
},
2026-08-15 13:50:24 +07:00
)
def stream(self, request: ModelRequestContract) -> Iterator[ModelChunk]:
config = self.resource.config
url = str(config.get("endpoint_url", "http://localhost:8000/v1/chat/completions"))
body = {
"model": config.get("model", self.resource.name),
"messages": [{"role": "user", "content": request.prompt}],
"max_tokens": request.token_budget,
"temperature": config.get("temperature", 0),
"stream": True,
"stream_options": {"include_usage": True},
}
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",
)
emitted = False
usage: dict[str, object] = {}
with urllib.request.urlopen(
http_request, timeout=int(config.get("timeout_seconds", 240))
) as response:
for raw_line in response:
line = raw_line.decode("utf-8").strip()
if not line.startswith("data:"):
continue
payload = line[5:].strip()
if payload == "[DONE]":
break
try:
event = json.loads(payload)
except json.JSONDecodeError:
continue
if event.get("usage"):
usage = event["usage"]
yield ModelChunk("", {"provider": self.provider_name, "usage": usage})
choices = event.get("choices") or []
if not choices:
continue
content = (choices[0].get("delta") or {}).get("content")
if content:
emitted = True
yield ModelChunk(
content,
{"provider": self.provider_name, "usage": usage},
)
if not emitted:
raise ProviderError("Qwen streaming provider response missing content")
2026-08-15 13:50:24 +07:00
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] = {}
opencode_resources = Resource.objects.filter(is_active=True, provider="opencode")
2026-08-15 13:50:24 +07:00
qwen = Resource.objects.filter(is_active=True, provider="local_inference").first()
for resource in opencode_resources:
model_key = str(resource.config.get("model_key") or "").strip().lower()
if model_key:
providers[model_key] = SolProvider(resource)
if "sol" not in providers:
sol = opencode_resources.first()
if sol is not None:
providers["sol"] = SolProvider(sol)
2026-08-15 13:50:24 +07:00
if qwen is not None:
providers["qwen"] = QwenProvider(qwen)
return providers