mirror of
https://github.com/Routstr/routstr-core.git
synced 2026-10-05 20:28:23 +00:00
feat: expose per-model reasoning effort on /v1/models
Keep upstream reasoning metadata (supported_efforts, default_effort, mandatory) on the catalog instead of dropping it on ingest, and map client reasoning_effort / reasoning.effort / thinking onto each model's allowlist before forwarding so unsupported levels are not sent upstream.
This commit is contained in:
+59
-18
@@ -5,7 +5,7 @@ import random
|
||||
import httpx
|
||||
from fastapi import APIRouter, Depends, HTTPException, Request
|
||||
from pydantic import BaseModel as V2BaseModel
|
||||
from pydantic.v1 import BaseModel
|
||||
from pydantic.v1 import BaseModel, validator
|
||||
from sqlmodel.ext.asyncio.session import AsyncSession
|
||||
|
||||
from ..core.db import ModelRow, UpstreamProviderRow, get_session
|
||||
@@ -85,6 +85,30 @@ class TopProvider(BaseModel):
|
||||
is_moderated: bool | None = None
|
||||
|
||||
|
||||
class Reasoning(BaseModel):
|
||||
"""Per-model reasoning-effort metadata, matching OpenRouter's shape."""
|
||||
|
||||
mandatory: bool | None = None
|
||||
default_enabled: bool | None = None
|
||||
supported_efforts: list[str] | None = None
|
||||
default_effort: str | None = None
|
||||
supports_max_tokens: bool | None = None
|
||||
|
||||
class Config:
|
||||
extra = "ignore"
|
||||
|
||||
def is_empty(self) -> bool:
|
||||
return not any(
|
||||
(
|
||||
self.mandatory is not None,
|
||||
self.default_enabled is not None,
|
||||
self.supported_efforts,
|
||||
self.default_effort,
|
||||
self.supports_max_tokens is not None,
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
class Model(BaseModel):
|
||||
id: str
|
||||
name: str
|
||||
@@ -101,10 +125,43 @@ class Model(BaseModel):
|
||||
canonical_slug: str | None = None
|
||||
alias_ids: list[str] | None = None
|
||||
forwarded_model_id: str | None = None
|
||||
reasoning: Reasoning | None = None
|
||||
|
||||
class Config:
|
||||
extra = "ignore"
|
||||
|
||||
def __hash__(self) -> int:
|
||||
return hash(self.id)
|
||||
|
||||
@validator("reasoning", pre=True)
|
||||
def _coerce_reasoning(cls, value: object) -> object:
|
||||
if value is None or value is False:
|
||||
return None
|
||||
if isinstance(value, Reasoning):
|
||||
return None if value.is_empty() else value
|
||||
if not isinstance(value, dict) or not value:
|
||||
return None
|
||||
try:
|
||||
parsed = Reasoning.parse_obj(value)
|
||||
except Exception:
|
||||
return None
|
||||
return None if parsed.is_empty() else parsed
|
||||
|
||||
def dict(self, **kwargs: object) -> dict:
|
||||
# Non-reasoning models omit the field entirely so the catalog stays
|
||||
# additive: existing clients never see a new null key.
|
||||
data = super().dict(**kwargs) # type: ignore[arg-type]
|
||||
reasoning = data.get("reasoning")
|
||||
if not reasoning:
|
||||
data.pop("reasoning", None)
|
||||
elif isinstance(reasoning, dict):
|
||||
cleaned = {k: v for k, v in reasoning.items() if v is not None}
|
||||
if cleaned:
|
||||
data["reasoning"] = cleaned
|
||||
else:
|
||||
data.pop("reasoning", None)
|
||||
return data
|
||||
|
||||
|
||||
def litellm_cost_entry(model_id: str) -> dict | None:
|
||||
"""Look up ``model_id`` in litellm's bundled cost map.
|
||||
@@ -474,23 +531,7 @@ def _update_model_sats_pricing(model: Model, sats_to_usd: float) -> Model:
|
||||
if (sats.max_cost or 0.0) < min_req_sats:
|
||||
sats.max_cost = min_req_sats
|
||||
|
||||
return Model(
|
||||
id=model.id,
|
||||
name=model.name,
|
||||
created=model.created,
|
||||
description=model.description,
|
||||
context_length=model.context_length,
|
||||
architecture=model.architecture,
|
||||
pricing=model.pricing,
|
||||
sats_pricing=sats,
|
||||
per_request_limits=model.per_request_limits,
|
||||
top_provider=model.top_provider,
|
||||
enabled=model.enabled,
|
||||
upstream_provider_id=model.upstream_provider_id,
|
||||
canonical_slug=model.canonical_slug,
|
||||
alias_ids=model.alias_ids,
|
||||
forwarded_model_id=model.forwarded_model_id,
|
||||
)
|
||||
return model.copy(update={"sats_pricing": sats})
|
||||
except Exception as e:
|
||||
logger.error(
|
||||
"Failed to update sats pricing for model",
|
||||
|
||||
@@ -64,6 +64,7 @@ from .cache_breakpoints import (
|
||||
from .count_tokens import MissingUsageEstimator, count_tokens_locally
|
||||
from .litellm_routing import detect_litellm_prefix
|
||||
from .rate_limit import UPSTREAM_RATE_LIMIT, classify_rate_limit
|
||||
from .reasoning_effort import apply_reasoning_effort
|
||||
|
||||
if typing.TYPE_CHECKING:
|
||||
from .ehbp import ConfidentialInferenceProfile, EHBPForwardingTarget
|
||||
@@ -710,8 +711,7 @@ class BaseUpstreamProvider:
|
||||
transformed_model = self.transform_model_name(original_model)
|
||||
data["input"]["model"] = transformed_model
|
||||
|
||||
# Ensure proper Responses API structure
|
||||
# Add any Responses-specific transformations here
|
||||
apply_reasoning_effort(data, model_obj)
|
||||
|
||||
return json.dumps(data).encode()
|
||||
except Exception as e:
|
||||
@@ -825,6 +825,9 @@ class BaseUpstreamProvider:
|
||||
if inject_anthropic_cache_breakpoints(data):
|
||||
changed = True
|
||||
|
||||
if apply_reasoning_effort(data, model_obj):
|
||||
changed = True
|
||||
|
||||
if changed:
|
||||
return json.dumps(data).encode()
|
||||
return body
|
||||
@@ -5422,22 +5425,8 @@ class BaseUpstreamProvider:
|
||||
{k: v * self.provider_fee for k, v in base_pricing.dict().items()}
|
||||
)
|
||||
|
||||
temp_model = Model(
|
||||
id=model.id,
|
||||
name=model.name,
|
||||
created=model.created,
|
||||
description=model.description,
|
||||
context_length=model.context_length,
|
||||
architecture=model.architecture,
|
||||
pricing=adjusted_pricing,
|
||||
sats_pricing=None,
|
||||
per_request_limits=model.per_request_limits,
|
||||
top_provider=model.top_provider,
|
||||
enabled=model.enabled,
|
||||
upstream_provider_id=model.upstream_provider_id,
|
||||
canonical_slug=model.canonical_slug,
|
||||
alias_ids=model.alias_ids,
|
||||
forwarded_model_id=model.forwarded_model_id,
|
||||
temp_model = model.copy(
|
||||
update={"pricing": adjusted_pricing, "sats_pricing": None}
|
||||
)
|
||||
|
||||
(
|
||||
@@ -5446,23 +5435,7 @@ class BaseUpstreamProvider:
|
||||
adjusted_pricing.max_cost,
|
||||
) = _calculate_usd_max_costs(temp_model)
|
||||
|
||||
return Model(
|
||||
id=model.id,
|
||||
name=model.name,
|
||||
created=model.created,
|
||||
description=model.description,
|
||||
context_length=model.context_length,
|
||||
architecture=model.architecture,
|
||||
pricing=adjusted_pricing,
|
||||
sats_pricing=model.sats_pricing,
|
||||
per_request_limits=model.per_request_limits,
|
||||
top_provider=model.top_provider,
|
||||
enabled=model.enabled,
|
||||
upstream_provider_id=model.upstream_provider_id,
|
||||
canonical_slug=model.canonical_slug,
|
||||
alias_ids=model.alias_ids,
|
||||
forwarded_model_id=model.forwarded_model_id,
|
||||
)
|
||||
return model.copy(update={"pricing": adjusted_pricing})
|
||||
|
||||
async def fetch_models(self) -> list[Model]:
|
||||
"""Fetch available models from upstream API and update cache.
|
||||
|
||||
@@ -32,6 +32,7 @@ from ..core.exceptions import UpstreamError
|
||||
from ..core.redaction import redact_org_ids
|
||||
from ..payment.models import Model
|
||||
from .rate_limit import classify_rate_limit
|
||||
from .reasoning_effort import adapt_messages_body_for_litellm
|
||||
|
||||
logger = get_logger(__name__)
|
||||
|
||||
@@ -71,6 +72,8 @@ ALLOWED_MESSAGES_REQUEST_FIELDS: frozenset[str] = frozenset(
|
||||
"tools",
|
||||
"tool_choice",
|
||||
"metadata",
|
||||
# OpenAI-shaped effort after thinking is lifted off Anthropic bodies.
|
||||
"reasoning_effort",
|
||||
}
|
||||
)
|
||||
|
||||
@@ -131,9 +134,7 @@ def parse_sse_blocks(buffer: bytes) -> tuple[list[dict], bytes]:
|
||||
return events, buffer
|
||||
|
||||
|
||||
def events_from_chunk(
|
||||
chunk: object, sse_buffer: bytes
|
||||
) -> tuple[list[dict], bytes]:
|
||||
def events_from_chunk(chunk: object, sse_buffer: bytes) -> tuple[list[dict], bytes]:
|
||||
"""Normalize a stream chunk into one or more event dicts.
|
||||
|
||||
``litellm.anthropic.messages.acreate(stream=True)`` yields raw SSE
|
||||
@@ -224,9 +225,7 @@ async def aggregate_anthropic_events_to_message(
|
||||
raw_json = partial_json.pop(idx, None)
|
||||
if raw_json is not None and idx < len(blocks):
|
||||
try:
|
||||
blocks[idx]["input"] = (
|
||||
json.loads(raw_json) if raw_json else {}
|
||||
)
|
||||
blocks[idx]["input"] = json.loads(raw_json) if raw_json else {}
|
||||
except json.JSONDecodeError:
|
||||
blocks[idx]["input"] = raw_json
|
||||
elif etype == "message_delta":
|
||||
@@ -468,9 +467,7 @@ async def dispatch_anthropic_messages(
|
||||
on bad input or upstream failure.
|
||||
"""
|
||||
if not request_body:
|
||||
raise UpstreamError(
|
||||
"Missing request body for /v1/messages", status_code=400
|
||||
)
|
||||
raise UpstreamError("Missing request body for /v1/messages", status_code=400)
|
||||
|
||||
try:
|
||||
body: dict = json.loads(request_body)
|
||||
@@ -489,6 +486,8 @@ async def dispatch_anthropic_messages(
|
||||
client_stream = bool(body.pop("stream", False))
|
||||
upstream_stream = True
|
||||
|
||||
adapt_messages_body_for_litellm(body, model_obj)
|
||||
|
||||
# Forward only allowlisted Anthropic Messages request fields. Any
|
||||
# other client-supplied key is dropped so it cannot leak into the
|
||||
# upstream request. See ALLOWED_MESSAGES_REQUEST_FIELDS.
|
||||
@@ -498,9 +497,7 @@ async def dispatch_anthropic_messages(
|
||||
"Dropped non-forwardable fields before litellm dispatch",
|
||||
extra={"dropped_keys": dropped},
|
||||
)
|
||||
body = {
|
||||
k: v for k, v in body.items() if k in ALLOWED_MESSAGES_REQUEST_FIELDS
|
||||
}
|
||||
body = {k: v for k, v in body.items() if k in ALLOWED_MESSAGES_REQUEST_FIELDS}
|
||||
|
||||
# Convention: `model.id` is the canonical upstream model name;
|
||||
# `forwarded_model_id` is the public alias the internal API exposes
|
||||
|
||||
@@ -66,9 +66,7 @@ class OllamaUpstreamProvider(BaseUpstreamProvider):
|
||||
"""Strip 'ollama/' prefix for Ollama API compatibility."""
|
||||
return model_id.removeprefix("ollama/")
|
||||
|
||||
def get_request_base_url(
|
||||
self, path: str, model_obj: Model | None = None
|
||||
) -> str:
|
||||
def get_request_base_url(self, path: str, model_obj: Model | None = None) -> str:
|
||||
"""Route proxy traffic through Ollama's OpenAI-compatible /v1 endpoint."""
|
||||
return f"{self.base_url.rstrip('/')}/v1"
|
||||
|
||||
@@ -185,7 +183,9 @@ class OllamaUpstreamProvider(BaseUpstreamProvider):
|
||||
except Exception:
|
||||
self._models_cache = models_with_fees
|
||||
|
||||
self._models_by_id = {m.forwarded_model_id or m.id: m for m in self._models_cache}
|
||||
self._models_by_id = {
|
||||
m.forwarded_model_id or m.id: m for m in self._models_cache
|
||||
}
|
||||
logger.info(
|
||||
f"Refreshed models cache for {self.base_url}",
|
||||
extra={"model_count": len(models)},
|
||||
@@ -224,26 +224,14 @@ class OllamaUpstreamProvider(BaseUpstreamProvider):
|
||||
Returns:
|
||||
Model with provider fee applied to pricing and max costs calculated
|
||||
"""
|
||||
from ..payment.models import Model, Pricing, _calculate_usd_max_costs
|
||||
from ..payment.models import Pricing, _calculate_usd_max_costs
|
||||
|
||||
adjusted_pricing = Pricing.parse_obj(
|
||||
{k: v * self.provider_fee for k, v in model.pricing.dict().items()}
|
||||
)
|
||||
|
||||
temp_model = Model(
|
||||
id=model.id,
|
||||
name=model.name,
|
||||
created=model.created,
|
||||
description=model.description,
|
||||
context_length=model.context_length,
|
||||
architecture=model.architecture,
|
||||
pricing=adjusted_pricing,
|
||||
sats_pricing=None,
|
||||
per_request_limits=model.per_request_limits,
|
||||
top_provider=model.top_provider,
|
||||
enabled=model.enabled,
|
||||
upstream_provider_id=model.upstream_provider_id,
|
||||
canonical_slug=model.canonical_slug,
|
||||
temp_model = model.copy(
|
||||
update={"pricing": adjusted_pricing, "sats_pricing": None}
|
||||
)
|
||||
|
||||
(
|
||||
@@ -252,18 +240,4 @@ class OllamaUpstreamProvider(BaseUpstreamProvider):
|
||||
adjusted_pricing.max_cost,
|
||||
) = _calculate_usd_max_costs(temp_model)
|
||||
|
||||
return Model(
|
||||
id=model.id,
|
||||
name=model.name,
|
||||
created=model.created,
|
||||
description=model.description,
|
||||
context_length=model.context_length,
|
||||
architecture=model.architecture,
|
||||
pricing=adjusted_pricing,
|
||||
sats_pricing=model.sats_pricing,
|
||||
per_request_limits=model.per_request_limits,
|
||||
top_provider=model.top_provider,
|
||||
enabled=model.enabled,
|
||||
upstream_provider_id=model.upstream_provider_id,
|
||||
canonical_slug=model.canonical_slug,
|
||||
)
|
||||
return model.copy(update={"pricing": adjusted_pricing})
|
||||
|
||||
@@ -0,0 +1,197 @@
|
||||
"""Map client reasoning/thinking effort onto a model's allowlist.
|
||||
|
||||
OpenRouter (and some other catalogs) publish per-model reasoning metadata:
|
||||
which effort levels are legal, the default, and whether reasoning is
|
||||
mandatory. Clients still send the generic OpenAI / Anthropic shapes
|
||||
(``reasoning_effort``, ``reasoning.effort``, ``thinking``). This module
|
||||
normalizes those into a supported effort and writes the fields the
|
||||
upstream actually accepts, instead of dropping the parameter or
|
||||
forwarding a value the model rejects.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
from ..payment.models import Model, Reasoning
|
||||
|
||||
# Highest first. Unknown values are treated as unranked.
|
||||
EFFORT_RANK: tuple[str, ...] = (
|
||||
"max",
|
||||
"xhigh",
|
||||
"high",
|
||||
"medium",
|
||||
"low",
|
||||
"minimal",
|
||||
"none",
|
||||
)
|
||||
_RANK_INDEX: dict[str, int] = {name: i for i, name in enumerate(EFFORT_RANK)}
|
||||
|
||||
_REASONING_KEYS = ("reasoning", "reasoning_effort", "thinking")
|
||||
|
||||
|
||||
def _normalize_effort(value: object) -> str | None:
|
||||
if not isinstance(value, str):
|
||||
return None
|
||||
cleaned = value.strip().lower()
|
||||
return cleaned or None
|
||||
|
||||
|
||||
def closest_supported_effort(
|
||||
requested: str | None,
|
||||
supported: list[str],
|
||||
*,
|
||||
default_effort: str | None = None,
|
||||
mandatory: bool = False,
|
||||
) -> str | None:
|
||||
"""Pick a legal effort for ``requested``.
|
||||
|
||||
Exact match wins. Otherwise the nearest rank in ``EFFORT_RANK`` is
|
||||
used (preferring the higher neighbour on a tie). ``none`` is rejected
|
||||
when ``mandatory`` is set. Missing / unmapped requests fall back to
|
||||
``default_effort``, then the highest remaining supported level.
|
||||
"""
|
||||
allowed_efforts: list[str] = [
|
||||
normalized
|
||||
for item in supported
|
||||
if (normalized := _normalize_effort(item)) is not None
|
||||
]
|
||||
if mandatory:
|
||||
allowed_efforts = [item for item in allowed_efforts if item != "none"]
|
||||
if not allowed_efforts:
|
||||
if mandatory:
|
||||
return _normalize_effort(default_effort)
|
||||
return _normalize_effort(requested) or _normalize_effort(default_effort)
|
||||
|
||||
default = _normalize_effort(default_effort)
|
||||
if default not in allowed_efforts:
|
||||
default = allowed_efforts[0]
|
||||
|
||||
requested_norm = _normalize_effort(requested)
|
||||
if requested_norm is None or (requested_norm == "none" and mandatory):
|
||||
return default
|
||||
|
||||
if requested_norm in allowed_efforts:
|
||||
return requested_norm
|
||||
|
||||
if requested_norm not in _RANK_INDEX:
|
||||
return default
|
||||
|
||||
target = _RANK_INDEX[requested_norm]
|
||||
return min(
|
||||
allowed_efforts,
|
||||
key=lambda effort: (
|
||||
abs(_RANK_INDEX.get(effort, 10_000) - target),
|
||||
_RANK_INDEX.get(effort, 10_000),
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
def resolve_effort(requested: str | None, reasoning: Reasoning | None) -> str | None:
|
||||
"""Map ``requested`` through ``reasoning`` metadata when present."""
|
||||
if reasoning is None:
|
||||
return _normalize_effort(requested)
|
||||
supported = reasoning.supported_efforts or []
|
||||
return closest_supported_effort(
|
||||
requested,
|
||||
supported,
|
||||
default_effort=reasoning.default_effort,
|
||||
mandatory=bool(reasoning.mandatory),
|
||||
)
|
||||
|
||||
|
||||
def _effort_from_thinking(thinking: object) -> str | None:
|
||||
if not isinstance(thinking, dict):
|
||||
return None
|
||||
effort = _normalize_effort(thinking.get("effort"))
|
||||
if effort:
|
||||
return effort
|
||||
thinking_type = _normalize_effort(thinking.get("type"))
|
||||
if thinking_type in {"disabled", "none"}:
|
||||
return "none"
|
||||
return None
|
||||
|
||||
|
||||
def extract_requested_effort(data: dict[str, Any]) -> str | None:
|
||||
"""Best-effort effort string from the OpenAI / Anthropic request shapes."""
|
||||
if isinstance(data.get("reasoning"), dict):
|
||||
nested = _normalize_effort(data["reasoning"].get("effort"))
|
||||
if nested:
|
||||
return nested
|
||||
top_level = _normalize_effort(data.get("reasoning_effort"))
|
||||
if top_level:
|
||||
return top_level
|
||||
return _effort_from_thinking(data.get("thinking"))
|
||||
|
||||
|
||||
def _request_mentions_reasoning(data: dict[str, Any]) -> bool:
|
||||
return any(key in data for key in _REASONING_KEYS)
|
||||
|
||||
|
||||
def apply_reasoning_effort(
|
||||
data: dict[str, Any],
|
||||
model: Model,
|
||||
*,
|
||||
drop_thinking: bool = False,
|
||||
) -> bool:
|
||||
"""Rewrite ``data`` in place so effort matches the model allowlist.
|
||||
|
||||
Returns True when ``data`` changed. Leaves the body alone when the
|
||||
caller did not send a reasoning field and the model does not require
|
||||
one. Existing ``reasoning`` object keys (``max_tokens``, ``exclude``,
|
||||
``enabled``) are preserved; only ``effort`` is mapped.
|
||||
|
||||
``drop_thinking`` is for OpenAI-compatible backends that reject the
|
||||
Anthropic ``thinking`` object: the effort is lifted onto
|
||||
``reasoning_effort`` / ``reasoning.effort`` and ``thinking`` is removed.
|
||||
"""
|
||||
if not isinstance(data, dict):
|
||||
return False
|
||||
|
||||
reasoning_meta = getattr(model, "reasoning", None)
|
||||
mentioned = _request_mentions_reasoning(data)
|
||||
if not mentioned and not (reasoning_meta and reasoning_meta.mandatory):
|
||||
return False
|
||||
|
||||
resolved = resolve_effort(extract_requested_effort(data), reasoning_meta)
|
||||
changed = False
|
||||
|
||||
if drop_thinking and "thinking" in data:
|
||||
data.pop("thinking", None)
|
||||
changed = True
|
||||
|
||||
if resolved is None:
|
||||
return changed
|
||||
|
||||
if "reasoning_effort" in data:
|
||||
if data.get("reasoning_effort") != resolved:
|
||||
data["reasoning_effort"] = resolved
|
||||
changed = True
|
||||
elif drop_thinking or (reasoning_meta and reasoning_meta.mandatory):
|
||||
# Invent the OpenAI-shaped field when we stripped Anthropic
|
||||
# ``thinking``, or when the model will reject a request with no
|
||||
# effort at all.
|
||||
if not isinstance(data.get("reasoning"), dict):
|
||||
data["reasoning_effort"] = resolved
|
||||
changed = True
|
||||
|
||||
existing = data.get("reasoning")
|
||||
if isinstance(existing, dict):
|
||||
if existing.get("effort") != resolved:
|
||||
data["reasoning"] = {**existing, "effort": resolved}
|
||||
changed = True
|
||||
elif reasoning_meta and reasoning_meta.mandatory and "reasoning_effort" not in data:
|
||||
data["reasoning"] = {"effort": resolved}
|
||||
changed = True
|
||||
|
||||
return changed
|
||||
|
||||
|
||||
def adapt_messages_body_for_litellm(data: dict[str, Any], model: Model) -> None:
|
||||
"""Convert Anthropic ``thinking`` into OpenAI-shaped effort for litellm.
|
||||
|
||||
Litellm's Anthropic-messages adapter talking to an OpenAI-compatible
|
||||
upstream will 400 on ``thinking``. Lift the effort onto
|
||||
``reasoning_effort`` and drop the Anthropic-only object.
|
||||
"""
|
||||
apply_reasoning_effort(data, model, drop_thinking=True)
|
||||
@@ -0,0 +1,261 @@
|
||||
"""Per-model reasoning-effort catalog metadata and request mapping."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import os
|
||||
from typing import Any
|
||||
|
||||
os.environ.setdefault("UPSTREAM_BASE_URL", "http://test")
|
||||
os.environ.setdefault("UPSTREAM_API_KEY", "test")
|
||||
os.environ.setdefault("LIGHTNING_ADDRESS", "test@stm.to")
|
||||
|
||||
from fastapi import FastAPI
|
||||
from fastapi.testclient import TestClient
|
||||
|
||||
from routstr.core.db import get_session
|
||||
from routstr.payment.models import (
|
||||
Architecture,
|
||||
Model,
|
||||
Pricing,
|
||||
Reasoning,
|
||||
models_router,
|
||||
)
|
||||
from routstr.upstream import GenericUpstreamProvider
|
||||
from routstr.upstream.reasoning_effort import (
|
||||
adapt_messages_body_for_litellm,
|
||||
apply_reasoning_effort,
|
||||
closest_supported_effort,
|
||||
extract_requested_effort,
|
||||
resolve_effort,
|
||||
)
|
||||
|
||||
|
||||
def _model(**kwargs: Any) -> Model:
|
||||
reasoning = kwargs.pop("reasoning", None)
|
||||
return Model(
|
||||
id=kwargs.get("id", "openai/gpt-5.6-sol"),
|
||||
name="test",
|
||||
created=0,
|
||||
description="",
|
||||
context_length=128000,
|
||||
architecture=Architecture(
|
||||
modality="text->text",
|
||||
input_modalities=["text"],
|
||||
output_modalities=["text"],
|
||||
tokenizer="x",
|
||||
instruct_type=None,
|
||||
),
|
||||
pricing=Pricing(prompt=0.0, completion=0.0),
|
||||
reasoning=reasoning,
|
||||
)
|
||||
|
||||
|
||||
SOL_REASONING = {
|
||||
"mandatory": False,
|
||||
"default_enabled": True,
|
||||
"supported_efforts": ["max", "xhigh", "high", "medium", "low", "none"],
|
||||
"default_effort": "medium",
|
||||
}
|
||||
|
||||
|
||||
def test_model_parses_openrouter_reasoning_object() -> None:
|
||||
model = Model(
|
||||
id="openai/gpt-5.6-sol",
|
||||
name="GPT",
|
||||
created=0,
|
||||
description="",
|
||||
context_length=1,
|
||||
architecture={
|
||||
"modality": "text",
|
||||
"input_modalities": ["text"],
|
||||
"output_modalities": ["text"],
|
||||
"tokenizer": "x",
|
||||
"instruct_type": None,
|
||||
},
|
||||
pricing={"prompt": 1e-6, "completion": 1e-6},
|
||||
reasoning=SOL_REASONING,
|
||||
extra_ignored_field="drop me",
|
||||
)
|
||||
assert model.reasoning is not None
|
||||
assert model.reasoning.supported_efforts == [
|
||||
"max",
|
||||
"xhigh",
|
||||
"high",
|
||||
"medium",
|
||||
"low",
|
||||
"none",
|
||||
]
|
||||
dumped = model.dict()
|
||||
assert dumped["reasoning"]["supported_efforts"][0] == "max"
|
||||
assert dumped["reasoning"]["default_effort"] == "medium"
|
||||
assert "extra_ignored_field" not in dumped
|
||||
|
||||
|
||||
def test_non_reasoning_models_omit_the_field() -> None:
|
||||
dumped = _model().dict()
|
||||
assert "reasoning" not in dumped
|
||||
|
||||
|
||||
def test_malformed_reasoning_is_dropped_not_fatal() -> None:
|
||||
model = _model(reasoning=["not", "a", "dict"])
|
||||
assert model.reasoning is None
|
||||
assert "reasoning" not in model.dict()
|
||||
|
||||
|
||||
def test_closest_effort_maps_minimal_to_low() -> None:
|
||||
assert (
|
||||
closest_supported_effort(
|
||||
"minimal",
|
||||
["max", "xhigh", "high", "medium", "low", "none"],
|
||||
default_effort="medium",
|
||||
)
|
||||
== "low"
|
||||
)
|
||||
|
||||
|
||||
def test_closest_effort_maps_max_when_missing() -> None:
|
||||
assert (
|
||||
closest_supported_effort(
|
||||
"max",
|
||||
["high", "medium", "low", "none"],
|
||||
default_effort="medium",
|
||||
)
|
||||
== "high"
|
||||
)
|
||||
|
||||
|
||||
def test_mandatory_rejects_none() -> None:
|
||||
assert (
|
||||
closest_supported_effort(
|
||||
"none",
|
||||
["max", "high", "medium", "low", "none"],
|
||||
default_effort="high",
|
||||
mandatory=True,
|
||||
)
|
||||
== "high"
|
||||
)
|
||||
|
||||
|
||||
def test_missing_request_uses_default() -> None:
|
||||
reasoning = Reasoning.parse_obj(SOL_REASONING)
|
||||
assert resolve_effort(None, reasoning) == "medium"
|
||||
|
||||
|
||||
def test_extract_prefers_nested_reasoning_effort() -> None:
|
||||
assert (
|
||||
extract_requested_effort(
|
||||
{"reasoning_effort": "low", "reasoning": {"effort": "high"}}
|
||||
)
|
||||
== "high"
|
||||
)
|
||||
|
||||
|
||||
def test_prepare_request_body_rewrites_unsupported_effort() -> None:
|
||||
provider = GenericUpstreamProvider(base_url="https://openrouter.ai/api/v1")
|
||||
model = _model(reasoning=SOL_REASONING)
|
||||
body = json.dumps(
|
||||
{
|
||||
"model": "openai/gpt-5.6-sol",
|
||||
"messages": [{"role": "user", "content": "hi"}],
|
||||
"reasoning_effort": "minimal",
|
||||
}
|
||||
).encode()
|
||||
out = provider.prepare_request_body(body, model)
|
||||
assert out is not None
|
||||
data = json.loads(out)
|
||||
assert data["reasoning_effort"] == "low"
|
||||
|
||||
|
||||
def test_prepare_request_body_rewrites_nested_effort() -> None:
|
||||
provider = GenericUpstreamProvider(base_url="https://openrouter.ai/api/v1")
|
||||
model = _model(reasoning=SOL_REASONING)
|
||||
body = json.dumps(
|
||||
{
|
||||
"model": "openai/gpt-5.6-sol",
|
||||
"messages": [{"role": "user", "content": "hi"}],
|
||||
"reasoning": {"effort": "minimal", "exclude": False},
|
||||
}
|
||||
).encode()
|
||||
out = provider.prepare_request_body(body, model)
|
||||
data = json.loads(out)
|
||||
assert data["reasoning"]["effort"] == "low"
|
||||
assert data["reasoning"]["exclude"] is False
|
||||
|
||||
|
||||
def test_prepare_request_body_leaves_plain_chat_alone() -> None:
|
||||
provider = GenericUpstreamProvider(base_url="https://openrouter.ai/api/v1")
|
||||
model = _model(reasoning=SOL_REASONING)
|
||||
payload = {
|
||||
"model": "openai/gpt-5.6-sol",
|
||||
"messages": [{"role": "user", "content": "hi"}],
|
||||
}
|
||||
body = json.dumps(payload).encode()
|
||||
out = provider.prepare_request_body(body, model)
|
||||
assert out == body
|
||||
|
||||
|
||||
def test_apply_injects_default_when_mandatory() -> None:
|
||||
data: dict[str, Any] = {
|
||||
"model": "anthropic/claude-fable-5.1",
|
||||
"messages": [{"role": "user", "content": "hi"}],
|
||||
}
|
||||
model = _model(
|
||||
id="anthropic/claude-fable-5.1",
|
||||
reasoning={
|
||||
"mandatory": True,
|
||||
"supported_efforts": ["max", "xhigh", "high", "medium", "low"],
|
||||
"default_effort": "high",
|
||||
},
|
||||
)
|
||||
assert apply_reasoning_effort(data, model) is True
|
||||
assert data["reasoning_effort"] == "high"
|
||||
|
||||
|
||||
def test_fee_apply_preserves_reasoning() -> None:
|
||||
provider = GenericUpstreamProvider(
|
||||
base_url="https://openrouter.ai/api/v1", provider_fee=1.1
|
||||
)
|
||||
model = _model(reasoning=SOL_REASONING)
|
||||
priced = provider._apply_provider_fee_to_model(model)
|
||||
assert priced.reasoning is not None
|
||||
assert priced.reasoning.supported_efforts == SOL_REASONING["supported_efforts"]
|
||||
|
||||
|
||||
def test_v1_models_includes_reasoning_and_omits_when_absent(
|
||||
monkeypatch: Any,
|
||||
) -> None:
|
||||
with_reasoning = _model(id="openai/gpt-5.6-sol", reasoning=SOL_REASONING)
|
||||
without = _model(id="openai/gpt-4o")
|
||||
unique = {"openai/gpt-5.6-sol": with_reasoning, "openai/gpt-4o": without}
|
||||
|
||||
import routstr.proxy as proxy
|
||||
|
||||
monkeypatch.setattr(proxy, "_unique_models", unique)
|
||||
app = FastAPI()
|
||||
app.include_router(models_router)
|
||||
app.dependency_overrides[get_session] = lambda: None
|
||||
response = TestClient(app).get("/v1/models")
|
||||
assert response.status_code == 200
|
||||
by_id = {row["id"]: row for row in response.json()["data"]}
|
||||
assert by_id["openai/gpt-5.6-sol"]["reasoning"]["supported_efforts"] == [
|
||||
"max",
|
||||
"xhigh",
|
||||
"high",
|
||||
"medium",
|
||||
"low",
|
||||
"none",
|
||||
]
|
||||
assert "reasoning" not in by_id["openai/gpt-4o"]
|
||||
|
||||
|
||||
def test_messages_thinking_becomes_reasoning_effort() -> None:
|
||||
model = _model(reasoning=SOL_REASONING)
|
||||
body: dict[str, Any] = {
|
||||
"messages": [{"role": "user", "content": "hi"}],
|
||||
"max_tokens": 16,
|
||||
"thinking": {"type": "enabled", "effort": "minimal"},
|
||||
}
|
||||
adapt_messages_body_for_litellm(body, model)
|
||||
assert "thinking" not in body
|
||||
assert body["reasoning_effort"] == "low"
|
||||
Reference in New Issue
Block a user