diff --git a/litellm/router_strategy/capability_router/policy.py b/litellm/router_strategy/capability_router/policy.py index d76a831e9d3..46303211704 100644 --- a/litellm/router_strategy/capability_router/policy.py +++ b/litellm/router_strategy/capability_router/policy.py @@ -43,7 +43,7 @@ def select_capability_model( verdict: CapabilityClassifierVerdict, estimated_costs: Mapping[str, float | None], ) -> CapabilityRoutingDecision: - """Choose the cheapest candidate at or above the configured probability.""" + """Choose the cheapest candidate above the configured probability.""" configured_models = tuple(candidate.model for candidate in config.candidates) scores = {candidate.model: candidate for candidate in verdict.candidates} if set(scores) != set(configured_models): @@ -55,7 +55,7 @@ def select_capability_model( p_solve=scores[model].p_solve, reason=scores[model].reason, estimated_cost=estimated_costs.get(model), - qualified=scores[model].p_solve >= config.probability_threshold, + qualified=scores[model].p_solve > config.probability_threshold, ) for model in configured_models ) diff --git a/tests/test_litellm/router_strategy/capability_router/test_capability_router.py b/tests/test_litellm/router_strategy/capability_router/test_capability_router.py index fd3fee10c5e..2c2879252cd 100644 --- a/tests/test_litellm/router_strategy/capability_router/test_capability_router.py +++ b/tests/test_litellm/router_strategy/capability_router/test_capability_router.py @@ -84,6 +84,22 @@ def test_policy_falls_back_if_no_model_qualifies_or_price_is_unknown() -> None: ) +def test_probability_must_be_strictly_above_threshold() -> None: + parsed = CapabilityRouterConfig.model_validate(config()) + verdict = CapabilityClassifierVerdict.model_validate( + { + "candidates": [ + {"model": "small", "p_solve": 0.7, "reason": "on the boundary"}, + {"model": "frontier", "p_solve": 0.7, "reason": "on the boundary"}, + ] + } + ) + + assert select_capability_model(parsed, verdict, {"small": 0.01, "frontier": 0.05}).reason == ( + "no_qualified_candidate" + ) + + def test_router_registers_capability_strategy() -> None: router = Router( model_list=[