Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
10 changes: 10 additions & 0 deletions README.md
Original file line number Diff line number Diff line change
Expand Up @@ -105,6 +105,16 @@ Interaction rules preserve their source-assigned `severity`. `DRUG_SAFETY_MIN_IN
- `POST /v1/hub/query-profiles/{profile}/roles/{role}/generate`: execute one
configured query role. Callers provide non-system messages and an optional
response format; callers cannot override the role model, prompt, or knobs.
Every request that reaches the model returns versioned request evidence with
the exact system and caller messages, selected profile/role/model, response
format and model settings, and an RFC 8785 canonical SHA-256 digest. When the
router supports it, the same evidence includes the exact rendered prompt and
token count against the model's advertised context window and the role's
configured output allowance. Missing rendering, counting, or window
information is reported with a stable reason and is never estimated. For an
otherwise valid request, only a known overflow blocks the model call; it
returns `422` with the same evidence. The existing `token_accounting` field
keeps its compatible successful shape.
- `POST /v1/hub/generate`: raw single-model compatibility endpoint for generic
consumers that own their own profile configuration. Catalyst does not use it.
- `GET /health`: service health, uptime, and process memory.
Expand Down
289 changes: 232 additions & 57 deletions server/generic_role.py
Original file line number Diff line number Diff line change
Expand Up @@ -17,21 +17,28 @@

from __future__ import annotations

from typing import Any, Dict, List, Optional
import hashlib
from copy import deepcopy
from typing import Any, Dict, List, Mapping, Optional

import httpx
import rfc8785
from fastapi import APIRouter, HTTPException
from pydantic import BaseModel, ConfigDict, Field

from . import team
from .config import llm_config
from .levels_loader import (
ModelNotFoundError,
catalyst_query_profile_metadata,
catalyst_query_profile_ids,
catalyst_query_profile_metadata,
get_catalyst_query_profile,
)
from .openai_compat import _backend_discovery_metadata, _served_backend_model_metadata
from .config import llm_config
from .openai_compat import (
ROUTER_PROBE_TIMEOUT_SECONDS,
_backend_discovery_metadata,
_served_backend_model_metadata,
)
from .prompt_loader import load_prompt

router = APIRouter()
Expand Down Expand Up @@ -65,9 +72,11 @@ class ProfileGenerateRequest(BaseModel):
class ProfileGenerateResponse(GenerateResponse):
profile_id: str
role: str
request_evidence: Dict[str, Any]
# Exact token evidence for the fully rendered request, counted with the
# model's own template and tokenizer before the call. None when the
# router could not count -- absence is honest, a character estimate is not.
# model's own template and tokenizer before the call. The compatible field
# is None unless all four legacy values are known; request_evidence records
# each missing fact and reason without substituting an estimate.
token_accounting: Optional[Dict[str, Any]] = None


Expand All @@ -76,69 +85,186 @@ def _backend_models() -> set[str] | None:
return None if discovered is None else set(discovered)


_CONTEXT_WINDOWS: Dict[str, int] = {}


def _context_window(model: str) -> Optional[int]:
"""The --ctx-size the router actually launched this model with."""
cached = _CONTEXT_WINDOWS.get(model)
if cached:
return cached
metadata = _served_backend_model_metadata() or {}
entry = metadata.get(model) or {}
args = list(((entry.get("status") or {}).get("args")) or [])
metadata = _served_backend_model_metadata()
if not isinstance(metadata, Mapping):
return None
entry = metadata.get(model)
if not isinstance(entry, Mapping):
return None
status = entry.get("status")
if not isinstance(status, Mapping):
return None
advertised_args = status.get("args")
if not isinstance(advertised_args, (list, tuple)):
return None
args = list(advertised_args)
for index, item in enumerate(args[:-1]):
if item in ("--ctx-size", "-c"):
try:
_CONTEXT_WINDOWS[model] = int(args[index + 1])
return _CONTEXT_WINDOWS[model]
value = int(args[index + 1])
if value <= 0:
return None
return value
except (TypeError, ValueError):
return None
return None


async def _token_accounting(
model: str, messages: List[Dict[str, Any]], max_tokens: Optional[int]
) -> Optional[Dict[str, Any]]:
"""Count the rendered request with the model's own template and tokenizer.
async def _prompt_measurement(
model: str, messages: List[Dict[str, Any]], output_reserve: int
) -> Dict[str, Any]:
"""Render and count the exact configured-role prompt when the router can.

/apply-template renders the exact prompt the model will consume --
special tokens included -- and /tokenize counts it with the model's own
vocabulary, so the number is the one the context window will see.
vocabulary. Missing router capabilities remain explicit; they are never
replaced with character or approximate token counts.
"""

window = _context_window(model)
if window is None or max_tokens is None:
return None
prompt: Optional[str] = None
prompt_tokens: Optional[int] = None
prompt_unavailable_reason: Optional[str] = None
count_unavailable_reason: Optional[str] = None
base = llm_config.base_url.rstrip("/")
headers = {}
if llm_config.api_key:
headers["Authorization"] = f"Bearer {llm_config.api_key}"
try:
async with httpx.AsyncClient(headers=headers, timeout=30.0) as client:
rendered = await client.post(
f"{base}/apply-template",
json={"model": model, "messages": messages},
)
rendered.raise_for_status()
prompt = rendered.json().get("prompt")
if not isinstance(prompt, str):
return None
tokenized = await client.post(
f"{base}/tokenize",
json={"model": model, "content": prompt},
)
tokenized.raise_for_status()
tokens = tokenized.json().get("tokens")
if not isinstance(tokens, list):
return None
async with httpx.AsyncClient(
headers=headers,
timeout=ROUTER_PROBE_TIMEOUT_SECONDS,
) as client:
try:
rendered = await client.post(
f"{base}/apply-template",
json={"model": model, "messages": messages},
)
rendered.raise_for_status()
candidate = rendered.json().get("prompt")
if isinstance(candidate, str):
prompt = candidate
else:
prompt_unavailable_reason = "prompt_rendering_unavailable"
except Exception:
prompt_unavailable_reason = "prompt_rendering_unavailable"

if prompt is None:
count_unavailable_reason = "rendered_prompt_unavailable"
else:
try:
tokenized = await client.post(
f"{base}/tokenize",
json={
"model": model,
"content": prompt,
# apply-template already emitted special tokens.
"add_special": False,
"parse_special": True,
},
)
tokenized.raise_for_status()
tokens = tokenized.json().get("tokens")
if isinstance(tokens, list):
prompt_tokens = len(tokens)
else:
count_unavailable_reason = "prompt_token_count_unavailable"
except Exception:
count_unavailable_reason = "prompt_token_count_unavailable"
except Exception:
return None
return {
if prompt is None:
prompt_unavailable_reason = "prompt_rendering_unavailable"
count_unavailable_reason = "rendered_prompt_unavailable"
elif prompt_tokens is None:
count_unavailable_reason = "prompt_token_count_unavailable"

prompt_evidence: Dict[str, Any] = {
"renderedPrompt": prompt,
"renderedPromptDigest": (
hashlib.sha256(prompt.encode("utf-8")).hexdigest()
if prompt is not None
else None
),
}
if prompt is None:
prompt_evidence["unavailableReason"] = (
prompt_unavailable_reason or "prompt_rendering_unavailable"
)

required_tokens = (
prompt_tokens + output_reserve if prompt_tokens is not None else None
)
fits = (
required_tokens <= window
if required_tokens is not None and window is not None
else None
)
token_evidence: Dict[str, Any] = {
"tokenizer": model,
"contextWindow": window,
"outputReserve": int(max_tokens),
"promptTokens": len(tokens),
"outputReserve": output_reserve,
"promptTokens": prompt_tokens,
"requiredTokens": required_tokens,
"fits": fits,
}
if window is None:
token_evidence["contextWindowUnavailableReason"] = "context_window_unavailable"
if prompt_tokens is None:
token_evidence["promptTokensUnavailableReason"] = (
count_unavailable_reason or "prompt_token_count_unavailable"
)
return {"prompt": prompt_evidence, "tokens": token_evidence}


def _request_evidence(
*,
profile_id: str,
role: str,
model: str,
messages: List[Dict[str, Any]],
response_format: Optional[Dict[str, Any]],
temperature: float,
dry_multiplier: float,
max_tokens: int,
measurement: Dict[str, Any],
) -> Dict[str, Any]:
"""Bind exact call inputs and router measurements with canonical JSON."""

request = {
"profileId": profile_id,
"role": role,
"model": model,
"messages": deepcopy(messages),
"responseFormat": deepcopy(response_format),
"config": {
"temperature": temperature,
"dryMultiplier": dry_multiplier,
"maxTokens": max_tokens,
},
}
return {
"contractVersion": "med-agent-hub.catalyst-role-request-evidence.v1",
"request": request,
"requestDigest": hashlib.sha256(rfc8785.dumps(request)).hexdigest(),
"prompt": deepcopy(measurement["prompt"]),
"tokens": deepcopy(measurement["tokens"]),
}


def _compatible_token_accounting(
measurement: Dict[str, Any],
) -> Optional[Dict[str, Any]]:
"""Keep the existing successful token_accounting response shape."""

tokens = measurement.get("tokens")
if not isinstance(tokens, dict):
return None
required = ("tokenizer", "contextWindow", "outputReserve", "promptTokens")
if any(tokens.get(key) is None for key in required):
return None
return {key: deepcopy(tokens[key]) for key in required}


def _profile_or_404(profile_id: str):
Expand Down Expand Up @@ -254,25 +380,74 @@ async def generate_query_role(
},
)
knobs = profile.knobs[role]
model = profile.models[role]
response_format = deepcopy(req.response_format)
temperature = float(knobs["temperature"])
dry_multiplier = float(knobs["dry"])
max_tokens = int(knobs["maxTokens"])
rendered_messages = [
{"role": "system", "content": load_prompt(profile.prompts[role])},
*req.messages,
*deepcopy(req.messages),
]
accounting = await _token_accounting(
profile.models[role], rendered_messages, int(knobs["maxTokens"])
)
content = await _chat_or_bad_gateway(
model=profile.models[role],
measurement = await _prompt_measurement(model, rendered_messages, max_tokens)
request_evidence = _request_evidence(
profile_id=profile_id,
role=role,
model=model,
messages=rendered_messages,
response_format=req.response_format,
temperature=float(knobs["temperature"]),
dry_multiplier=float(knobs["dry"]),
max_tokens=int(knobs["maxTokens"]),
response_format=response_format,
temperature=temperature,
dry_multiplier=dry_multiplier,
max_tokens=max_tokens,
measurement=measurement,
)
if measurement["tokens"].get("fits") is False:
raise HTTPException(
status_code=422,
detail={
"code": "context_window_exceeded",
"message": (
"The exact rendered role request plus its configured output "
"reserve exceeds the model context window."
),
"request_evidence": request_evidence,
},
)
try:
content = await _chat_or_bad_gateway(
model=model,
messages=rendered_messages,
response_format=response_format,
temperature=temperature,
dry_multiplier=dry_multiplier,
max_tokens=max_tokens,
)
except HTTPException as error:
raise HTTPException(
status_code=error.status_code,
detail={
"code": "model_request_failed",
"message": str(error.detail),
"request_evidence": request_evidence,
},
headers=error.headers,
) from error
except Exception as error:
raise HTTPException(
status_code=502,
detail={
"code": "model_request_failed",
"message": (
"The model backend did not return a usable assistant response."
),
"request_evidence": request_evidence,
},
) from error
return ProfileGenerateResponse(
profile_id=profile_id,
role=role,
model=profile.models[role],
model=model,
content=content,
token_accounting=accounting,
request_evidence=request_evidence,
token_accounting=_compatible_token_accounting(measurement),
)
Loading
Loading