mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-16 23:41:43 +00:00
fix(router): expose exact Fuse v2 forecast metadata
This commit is contained in:
parent
909a7cd515
commit
92bece2baa
4 changed files with 109 additions and 8 deletions
|
|
@ -103,7 +103,7 @@ from .config import (
|
|||
CustomDimension,
|
||||
TierDefinition,
|
||||
)
|
||||
from .llm_v2 import LLMV2TaskContext, LLMV2Verdict, llm_v2_response_format
|
||||
from .llm_v2 import LLM_V2_PROMPT_VERSION, LLMV2Decision, LLMV2TaskContext, LLMV2Verdict, llm_v2_response_format
|
||||
from .stall_detector import detect_stalled_task
|
||||
|
||||
if TYPE_CHECKING:
|
||||
|
|
@ -1017,16 +1017,41 @@ class ClassificationOutcome(NamedTuple):
|
|||
]
|
||||
classifier_cost: float | None = None
|
||||
capability_forecast: CapabilityClassifierForecast | None = None
|
||||
llm_v2_forecast: LLMV2Decision | None = None
|
||||
|
||||
|
||||
def _with_signal(outcome: ClassificationOutcome, signal: str | None) -> ClassificationOutcome:
|
||||
return outcome if signal is None else outcome._replace(signals=(*outcome.signals, signal))
|
||||
|
||||
|
||||
def _with_capability_forecast(
|
||||
def _with_llm_v2_forecast(
|
||||
decision: StandardLoggingRoutingDecision, forecast: LLMV2Decision
|
||||
) -> StandardLoggingRoutingDecision:
|
||||
"""Preserve full numeric precision for both solver forecasts and the applied policy."""
|
||||
enriched: Final[StandardLoggingRoutingDecision] = {
|
||||
**decision,
|
||||
"classifier_efficient_p_solve": forecast.verdict.forecasts.efficient.p_solve,
|
||||
"classifier_capable_p_solve": forecast.verdict.forecasts.capable.p_solve,
|
||||
"classifier_max_quality_gap": forecast.max_quality_gap,
|
||||
"classifier_prompt_version": LLM_V2_PROMPT_VERSION,
|
||||
}
|
||||
if forecast.calibration_version is None:
|
||||
return enriched
|
||||
calibrated: Final[StandardLoggingRoutingDecision] = {
|
||||
**enriched,
|
||||
"classifier_calibrated_efficient_p_solve": forecast.efficient,
|
||||
"classifier_calibrated_capable_p_solve": forecast.capable,
|
||||
"classifier_calibration_version": forecast.calibration_version,
|
||||
}
|
||||
return calibrated
|
||||
|
||||
|
||||
def _with_classifier_forecast(
|
||||
decision: StandardLoggingRoutingDecision, outcome: ClassificationOutcome
|
||||
) -> StandardLoggingRoutingDecision:
|
||||
"""Attach the validated capability verdict and applied threshold to its decision record."""
|
||||
"""Attach validated forecasts and their applied policy to the routing decision."""
|
||||
if outcome.llm_v2_forecast is not None:
|
||||
return _with_llm_v2_forecast(decision, outcome.llm_v2_forecast)
|
||||
forecast: Final = outcome.capability_forecast
|
||||
if forecast is None:
|
||||
return decision
|
||||
|
|
@ -2347,6 +2372,7 @@ class ComplexityRouter(CustomLogger):
|
|||
signals=decision.signals,
|
||||
cause="llm_v2_classifier",
|
||||
classifier_cost=classifier_cost,
|
||||
llm_v2_forecast=decision,
|
||||
)
|
||||
|
||||
async def _call_classifier_model(
|
||||
|
|
@ -4474,5 +4500,5 @@ class ComplexityRouter(CustomLogger):
|
|||
model=routed_model,
|
||||
messages=messages if has_original_messages else None,
|
||||
litellm_params=tier_litellm_params,
|
||||
routing_decision=_with_capability_forecast(routing_decision, outcome),
|
||||
routing_decision=_with_classifier_forecast(routing_decision, outcome),
|
||||
)
|
||||
|
|
|
|||
|
|
@ -2990,6 +2990,12 @@ class StandardLoggingRoutingDecision(TypedDict, total=False):
|
|||
classifier_p_solve: float # writable-ok: added only when a capability verdict is available
|
||||
classifier_calibrated_p_solve: ReadOnly[float]
|
||||
classifier_calibration_version: ReadOnly[str]
|
||||
classifier_efficient_p_solve: ReadOnly[float]
|
||||
classifier_capable_p_solve: ReadOnly[float]
|
||||
classifier_calibrated_efficient_p_solve: ReadOnly[float]
|
||||
classifier_calibrated_capable_p_solve: ReadOnly[float]
|
||||
classifier_max_quality_gap: ReadOnly[float]
|
||||
classifier_prompt_version: ReadOnly[str]
|
||||
classifier_threshold: float # writable-ok: added only when a capability verdict is available
|
||||
escalated: bool
|
||||
context_escalated: bool # writable-ok: Pydantic warns on ReadOnly TypedDict fields
|
||||
|
|
@ -3026,6 +3032,12 @@ DERIVED_ROUTING_DECISION_FIELDS: Final[frozenset[str]] = frozenset(
|
|||
"classifier_p_solve",
|
||||
"classifier_calibrated_p_solve",
|
||||
"classifier_calibration_version",
|
||||
"classifier_efficient_p_solve",
|
||||
"classifier_capable_p_solve",
|
||||
"classifier_calibrated_efficient_p_solve",
|
||||
"classifier_calibrated_capable_p_solve",
|
||||
"classifier_max_quality_gap",
|
||||
"classifier_prompt_version",
|
||||
"classifier_threshold",
|
||||
"escalated",
|
||||
"context_escalated",
|
||||
|
|
|
|||
|
|
@ -4,6 +4,7 @@ from typing import Final
|
|||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
import pytest
|
||||
import litellm
|
||||
from pydantic import ValidationError
|
||||
|
||||
from litellm import ModelResponse, Router
|
||||
|
|
@ -11,6 +12,7 @@ from litellm.caching.dual_cache import DualCache
|
|||
from litellm.router_strategy.complexity_router.complexity_router import ComplexityRouter
|
||||
from litellm.router_strategy.complexity_router.config import ComplexityRouterConfig, ComplexityTier
|
||||
from litellm.router_strategy.complexity_router.llm_v2 import (
|
||||
LLM_V2_PROMPT_VERSION,
|
||||
LLMV2Calibration,
|
||||
LLMV2Config,
|
||||
LLMV2ProbabilityCalibration,
|
||||
|
|
@ -219,14 +221,63 @@ async def test_json_object_mode_supplies_schema_in_prompt() -> None:
|
|||
assert '"required"' in sent["messages"][0]["content"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("calibrated", (False, True))
|
||||
async def test_routing_metadata_preserves_exact_forecasts_and_redaction(
|
||||
calibrated: bool, monkeypatch: pytest.MonkeyPatch
|
||||
) -> None:
|
||||
base: Final = _config().llm_v2_config
|
||||
assert base is not None
|
||||
calibration: Final = LLMV2Calibration(
|
||||
version="test-pair-v1",
|
||||
prompt_version=LLM_V2_PROMPT_VERSION,
|
||||
efficient=LLMV2ProbabilityCalibration(slope=0.2, intercept=-1.0),
|
||||
capable=LLMV2ProbabilityCalibration(slope=1.0, intercept=0.0),
|
||||
)
|
||||
policy: Final = base.model_copy(update={"calibration": calibration if calibrated else None})
|
||||
verdict: Final = _verdict(0.900000123, 0.920000321)
|
||||
router, _ = _router(verdict.model_dump_json(), _config(llm_v2_config=policy.model_dump()))
|
||||
result: Final = await router.async_pre_routing_hook(
|
||||
model="v2-router", messages=[{"role": "user", "content": "Fix nested behavior"}], request_kwargs={}
|
||||
)
|
||||
assert result is not None
|
||||
assert result.model == ("capable" if calibrated else "efficient")
|
||||
decision: Final = result.routing_decision
|
||||
assert decision is not None
|
||||
monkeypatch.setattr(litellm, "turn_off_message_logging", True)
|
||||
redacted: Final = Router._redact_prompt_text_if_needed(request_kwargs={}, routing_decision=decision)
|
||||
assert redacted is not None
|
||||
assert "signals" not in redacted
|
||||
for record in (decision, redacted):
|
||||
assert record["classifier_efficient_p_solve"] == 0.900000123
|
||||
assert record["classifier_capable_p_solve"] == 0.920000321
|
||||
assert record["classifier_max_quality_gap"] == 0.05
|
||||
assert record["classifier_prompt_version"] == LLM_V2_PROMPT_VERSION
|
||||
if calibrated:
|
||||
assert record["classifier_calibration_version"] == "test-pair-v1"
|
||||
assert record["classifier_calibrated_efficient_p_solve"] == calibration.efficient.calibrate(0.900000123)
|
||||
assert record["classifier_calibrated_capable_p_solve"] == calibration.capable.calibrate(0.920000321)
|
||||
else:
|
||||
assert "classifier_calibration_version" not in record
|
||||
assert "classifier_calibrated_efficient_p_solve" not in record
|
||||
assert "classifier_calibrated_capable_p_solve" not in record
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("content", ["", "not json", '{"tier":"SIMPLE"}', '{"forecasts":{}}'])
|
||||
async def test_invalid_output_falls_back_to_capable_and_preserves_paid_call_cost(content: str) -> None:
|
||||
router, client = _router(content)
|
||||
outcome: Final = await router.aclassify("hi")
|
||||
assert outcome.tier == ComplexityTier.REASONING
|
||||
assert outcome.cause == "llm_v2_fallback"
|
||||
assert outcome.classifier_cost == 0.001
|
||||
result: Final = await router.async_pre_routing_hook(
|
||||
model="v2-router", messages=[{"role": "user", "content": "hi"}], request_kwargs={}
|
||||
)
|
||||
assert result is not None and result.model == "capable"
|
||||
decision: Final = result.routing_decision
|
||||
assert decision is not None
|
||||
assert decision["cause"] == "llm_v2_fallback"
|
||||
assert decision["classifier_cost"] == 0.001
|
||||
assert "classifier_efficient_p_solve" not in decision
|
||||
assert "classifier_capable_p_solve" not in decision
|
||||
assert "classifier_prompt_version" not in decision
|
||||
client.acompletion.assert_awaited_once()
|
||||
|
||||
|
||||
|
|
|
|||
12
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
12
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
|
|
@ -37068,22 +37068,34 @@ export interface components {
|
|||
* @enum {string}
|
||||
*/
|
||||
cause?: "heuristic_scorer" | "heuristic_v2" | "reasoning_override" | "llm_classifier" | "capability_classifier" | "llm_v2_classifier" | "llm_v2_fallback" | "heuristic_first_short_circuit" | "hybrid_short_circuit" | "classifier_plugin" | "classifier_fallback" | "capability_classifier_fallback" | "default_model_fallback" | "literal_keyword_match" | "semantic_keyword_match" | "plan_mode" | "housekeeping" | "modality_escalation" | "modality_pin_override" | "health_failover" | "health_default_fallback" | "session_affinity_pin" | "session_affinity_escalation" | "user_turn_continuation" | "default_fallback" | "keyword" | "quality_tier" | "bandit";
|
||||
/** Classifier Calibrated Capable P Solve */
|
||||
classifier_calibrated_capable_p_solve?: number;
|
||||
/** Classifier Calibrated Efficient P Solve */
|
||||
classifier_calibrated_efficient_p_solve?: number;
|
||||
/** Classifier Calibrated P Solve */
|
||||
classifier_calibrated_p_solve?: number;
|
||||
/** Classifier Calibration Version */
|
||||
classifier_calibration_version?: string;
|
||||
/** Classifier Capability Boundary */
|
||||
classifier_capability_boundary?: string;
|
||||
/** Classifier Capable P Solve */
|
||||
classifier_capable_p_solve?: number;
|
||||
/** Classifier Cost */
|
||||
classifier_cost?: number;
|
||||
/** Classifier Crux */
|
||||
classifier_crux?: string;
|
||||
/** Classifier Efficient P Solve */
|
||||
classifier_efficient_p_solve?: number;
|
||||
/** Classifier Max Quality Gap */
|
||||
classifier_max_quality_gap?: number;
|
||||
/** Classifier Model */
|
||||
classifier_model?: string;
|
||||
/** Classifier P Solve */
|
||||
classifier_p_solve?: number;
|
||||
/** Classifier Primary Rule */
|
||||
classifier_primary_rule?: string;
|
||||
/** Classifier Prompt Version */
|
||||
classifier_prompt_version?: string;
|
||||
/** Classifier Threshold */
|
||||
classifier_threshold?: number;
|
||||
/** Context Escalated */
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue