diff --git a/README.md b/README.md index adad85e..6d0a890 100644 --- a/README.md +++ b/README.md @@ -181,7 +181,11 @@ These are real outputs from `examples/basic.py`. | `router.fitting(task, limits=None) -> list[ModelInfo]` | Returns the models that pass the limits, without calling Jev (free). | | `Limits(output_tokens=1024, max_cost_usd=None)` | The output size you expect and an optional cost cap for each call. | -Errors (all subclasses of `RouterError`): +`Limits` validates its arguments when created: `output_tokens` must be an integer of at least 1, +and `max_cost_usd` must be `None` (no cost cap) or a nonnegative `int` or `float` (zero is allowed). +Booleans and NaN are rejected. Invalid limits raise `ValueError` with the field name and required range. + +Routing errors (all subclasses of `RouterError`): - `NoModelFitsError`: no model passes the limits. The message gives the reason for each model. - `UnknownModelError`: a model id isn't on OpenRouter. - `RouterError`: no routing key, a network or HTTP failure, or an error returned by Jev (e.g. a rate limit). diff --git a/src/model_router/models.py b/src/model_router/models.py index 0ff700c..5e1f0da 100644 --- a/src/model_router/models.py +++ b/src/model_router/models.py @@ -36,3 +36,17 @@ def from_openrouter(cls, raw): class Limits: output_tokens: int = 1024 max_cost_usd: float | None = None + + def __post_init__(self): + if ( + isinstance(self.output_tokens, bool) + or not isinstance(self.output_tokens, int) + or self.output_tokens < 1 + ): + raise ValueError("output_tokens must be an integer >= 1") + if self.max_cost_usd is not None and ( + isinstance(self.max_cost_usd, bool) + or not isinstance(self.max_cost_usd, (int, float)) + or not self.max_cost_usd >= 0 + ): + raise ValueError("max_cost_usd must be None or a number >= 0") diff --git a/tests/test_models.py b/tests/test_models.py index f415f38..963c8df 100644 --- a/tests/test_models.py +++ b/tests/test_models.py @@ -56,6 +56,28 @@ class LimitsTest(unittest.TestCase): def test_defaults(self): self.assertEqual(Limits(), Limits(output_tokens=1024, max_cost_usd=None)) + def test_positive_integer_output_tokens_are_accepted(self): + for output_tokens in (1, 500, 1024): + with self.subTest(output_tokens=output_tokens): + self.assertEqual(Limits(output_tokens=output_tokens).output_tokens, output_tokens) + + def test_invalid_output_tokens_raise_value_error(self): + for output_tokens in (0, -5, 1.0, 1.5, "1", None, True, False): + with self.subTest(output_tokens=output_tokens): + with self.assertRaisesRegex(ValueError, "output_tokens must be an integer >= 1"): + Limits(output_tokens=output_tokens) + + def test_optional_nonnegative_cost_is_accepted(self): + for max_cost_usd in (None, 0, 0.0, 0.001, 1): + with self.subTest(max_cost_usd=max_cost_usd): + self.assertEqual(Limits(max_cost_usd=max_cost_usd).max_cost_usd, max_cost_usd) + + def test_invalid_cost_raises_value_error(self): + for max_cost_usd in (-1, -0.001, float("nan"), "0", True, False): + with self.subTest(max_cost_usd=max_cost_usd): + with self.assertRaisesRegex(ValueError, "max_cost_usd must be None or a number >= 0"): + Limits(max_cost_usd=max_cost_usd) + def test_is_immutable(self): with self.assertRaises(dataclasses.FrozenInstanceError): Limits().output_tokens = 5