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 from model_router.policy import model_for_purpose 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: return model_for_purpose(str(purpose)) 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 in {"sol", "terra", "luna"}: provider_names.add("opencode") for resource in candidates: if resource.config.get("model_key") == provider_key: return resource 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", ] )