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

158 lines
5.5 KiB
Python

from __future__ import annotations
import time
import uuid
from dataclasses import dataclass
from enum import StrEnum
from typing import Protocol
from django.utils import timezone
from control_plane.agents.models import AgentVersion
from control_plane.projects.models import Project
from control_plane.resources.models import ModelRequest, Resource
class ModelCapability(StrEnum):
PROJECT_BRAIN = "PROJECT_BRAIN"
PLANNING = "PLANNING"
ARCHAEOLOGY_INTERPRETATION = "ARCHAEOLOGY_INTERPRETATION"
CODING = "CODING"
REVIEW = "REVIEW"
REASONING = "REASONING"
@dataclass(frozen=True)
class ModelRequestContract:
purpose: str
prompt: str
model_hint: str | None = None
token_budget: int = 8000
project: Project | None = None
agent_version: AgentVersion | None = None
@dataclass(frozen=True)
class ModelResponseContract:
model: str
content: str
metadata: dict[str, object]
class ModelProvider(Protocol):
provider_name: str
def complete(self, request: ModelRequestContract) -> ModelResponseContract: ...
def health(self) -> str: ...
class ModelRouter:
def __init__(self, providers: dict[str, ModelProvider] | None = None, persist_requests: bool = False) -> None:
self.providers = providers or {}
self.persist_requests = persist_requests
def route(self, purpose: str) -> str:
normalized = purpose.upper()
if normalized in {
ModelCapability.PROJECT_BRAIN,
ModelCapability.PLANNING,
ModelCapability.ARCHAEOLOGY_INTERPRETATION,
"PLANNING",
"ARCHAEOLOGY_INTERPRETATION",
"AGENT_DESIGN",
"ESCALATION",
}:
return "sol"
return "qwen"
def complete(self, request: ModelRequestContract) -> ModelResponseContract:
provider_key = request.model_hint or self.route(request.purpose)
provider = self.providers.get(provider_key)
if provider is None:
raise RuntimeError(f"No model provider configured for {provider_key}")
model_resource = self._resource_for(provider_key, request.purpose)
record = self._start_record(request, provider, model_resource)
started = time.monotonic()
try:
response = provider.complete(request)
except Exception as exc:
if record is not None:
self._finish_record(record, "FAILED", None, started, failure_reason=str(exc))
raise
if record is not None:
self._finish_record(record, "COMPLETE", response, started)
return response
def health(self) -> dict[str, str]:
statuses: dict[str, str] = {}
for key, provider in self.providers.items():
try:
statuses[key] = provider.health()
except Exception:
statuses[key] = "UNAVAILABLE"
return statuses
def _resource_for(self, provider_key: str, purpose: str) -> Resource | None:
role = purpose.upper()
candidates = list(Resource.objects.filter(is_active=True))
provider_names = {provider_key}
if provider_key == "qwen":
provider_names.add("local_inference")
if provider_key == "sol":
provider_names.add("opencode")
for resource in candidates:
if resource.provider in provider_names and role in resource.roles:
return resource
for resource in candidates:
if resource.provider in provider_names:
return resource
return None
def _start_record(self, request: ModelRequestContract, provider: ModelProvider, resource: Resource | None) -> ModelRequest | None:
if not self.persist_requests:
return None
return ModelRequest.objects.create(
correlation_id=uuid.uuid4().hex,
project=request.project,
agent_version=request.agent_version,
logical_role=request.purpose.upper(),
model_resource=resource,
provider=provider.provider_name,
model=resource.name if resource else request.model_hint or self.route(request.purpose),
token_budget=request.token_budget,
status="IN_PROGRESS",
started_at=timezone.now(),
request={"prompt_chars": len(request.prompt), "contains_raw_prompt": False},
)
def _finish_record(
self,
record: ModelRequest,
status: str,
response: ModelResponseContract | None,
started: float,
*,
failure_reason: str = "",
) -> None:
record.status = status
record.ended_at = timezone.now()
record.latency_ms = int((time.monotonic() - started) * 1000)
record.failure_reason = failure_reason
if response is not None:
usage = response.metadata.get("usage", {}) if isinstance(response.metadata, dict) else {}
record.prompt_tokens = usage.get("prompt_tokens") if isinstance(usage, dict) else None
record.completion_tokens = usage.get("completion_tokens") if isinstance(usage, dict) else None
record.response = {"content_chars": len(response.content), "model": response.model, "metadata_keys": sorted(response.metadata.keys())}
record.save(
update_fields=[
"status",
"ended_at",
"latency_ms",
"failure_reason",
"prompt_tokens",
"completion_tokens",
"response",
"updated_at",
]
)