158 lines
5.5 KiB
Python
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",
|
|
]
|
|
)
|