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
2 changes: 1 addition & 1 deletion README.md
Original file line number Diff line number Diff line change
Expand Up @@ -174,7 +174,7 @@ These are real outputs from `examples/basic.py`.

| Call | What it does |
|---|---|
| `Router(*, jev_api_key=None, openrouter_api_key=None, providers=None, models=None, limits=Limits(), models_per_provider=None)` | All arguments are keyword-only. Loads live prices, context sizes and output limits (cached for 24h, no key needed). Raises `UnknownModelError` for unknown model ids or providers. |
| `Router(*, jev_api_key=None, openrouter_api_key=None, providers=None, models=None, limits=Limits(), models_per_provider=None, timeout=60)` | All arguments are keyword-only. Loads live prices, context sizes and output limits (cached for 24h, no key needed). `timeout` is the HTTP timeout in seconds for catalog downloads and Jev/OpenRouter routing calls (e.g. `timeout=5`); it is passed to `urllib.request.urlopen`, not a total routing deadline. Raises `UnknownModelError` for unknown model ids or providers. |
| `refresh_catalog()` | Clears the cached model list, so the next `Router` downloads fresh prices. |
| `router.api_key_for(model_id) -> str \| None` | The key you passed in `providers={...}` for this model's provider. |
| `router.route(task, limits=None) -> str` | Returns the best model id for `task`. |
Expand Down
4 changes: 2 additions & 2 deletions src/model_router/catalog.py
Original file line number Diff line number Diff line change
Expand Up @@ -9,7 +9,7 @@
_cache = {"at": 0.0, "data": None}


def fetch_catalog(max_age=CATALOG_TTL_SECONDS):
def fetch_catalog(max_age=CATALOG_TTL_SECONDS, *, timeout=60):
"""Live prices, context and output limits for every text model on OpenRouter, keyed by id.

Public endpoint, no key needed. Fetched once and shared by every Router in the process,
Expand All @@ -19,7 +19,7 @@ def fetch_catalog(max_age=CATALOG_TTL_SECONDS):
if _cache["data"] is None or time.time() - _cache["at"] > max_age:
_cache["data"] = {
m["id"]: ModelInfo.from_openrouter(m)
for m in request_json(MODELS_URL)["data"]
for m in request_json(MODELS_URL, timeout=timeout)["data"]
if "text" in ((m.get("architecture") or {}).get("output_modalities") or ["text"])
}
_cache["at"] = time.time()
Expand Down
6 changes: 4 additions & 2 deletions src/model_router/jev.py
Original file line number Diff line number Diff line change
Expand Up @@ -51,7 +51,7 @@ def _fit(build, n_candidates):
)


def choose_via_jev(api_key, task, in_tokens, candidates, limits):
def choose_via_jev(api_key, task, in_tokens, candidates, limits, *, timeout=60):
"""jevai.org model-route preset. Returns the chosen model id."""
res = request_json(
JEV_ROUTE_URL,
Expand All @@ -69,14 +69,15 @@ def choose_via_jev(api_key, task, in_tokens, candidates, limits):
},
len(candidates),
),
timeout=timeout,
)
decision = (res.get("data") or {}).get("decision")
if res.get("code") != 0 or decision not in {m.id for m in candidates}:
raise RouterError(f"Unexpected Jev response: {res}")
return decision


def choose_via_openrouter(api_key, task, in_tokens, candidates, limits):
def choose_via_openrouter(api_key, task, in_tokens, candidates, limits, *, timeout=60):
"""OpenRouter decisions API running Jev. Returns the chosen model id."""
# Aliases keep criteria keys plain; model ids contain "/" and ".".
alias = {f"m{i}": m for i, m in enumerate(candidates)}
Expand All @@ -97,6 +98,7 @@ def choose_via_openrouter(api_key, task, in_tokens, candidates, limits):
},
len(candidates),
),
timeout=timeout,
)
picked = ((res.get("answers") or {}).get("model") or {}).get("choice")
if picked not in alias:
Expand Down
8 changes: 5 additions & 3 deletions src/model_router/router.py
Original file line number Diff line number Diff line change
Expand Up @@ -22,6 +22,7 @@ def __init__(
models=None,
limits=Limits(),
models_per_provider=None,
timeout=60,
):
"""
Routing backend (one is required; jev_api_key wins if both are given):
Expand All @@ -33,18 +34,19 @@ def __init__(
models exact OpenRouter-style ids, e.g. ["anthropic/claude-opus-5.5"]; overrides the auto-pick
"""
if jev_api_key:
self._choose = partial(jev.choose_via_jev, jev_api_key)
self._choose = partial(jev.choose_via_jev, jev_api_key, timeout=timeout)
elif openrouter_api_key:
self._choose = partial(jev.choose_via_openrouter, openrouter_api_key)
self._choose = partial(jev.choose_via_openrouter, openrouter_api_key, timeout=timeout)
else:
raise RouterError("Pass jev_api_key or openrouter_api_key")
if not providers and not models:
raise RouterError("Pass providers (names or {name: api_key}) and/or models")

self.provider_keys = dict(providers) if isinstance(providers, dict) else {}
self.limits = limits
self.timeout = timeout
# Live prices/context/limits, fetched once and cached (see catalog.fetch_catalog).
catalog = fetch_catalog()
catalog = fetch_catalog(timeout=self.timeout)

if models:
models = list(dict.fromkeys(models))
Expand Down
42 changes: 41 additions & 1 deletion tests/test_router.py
Original file line number Diff line number Diff line change
@@ -1,7 +1,8 @@
import asyncio
import io
import json
import unittest
from unittest.mock import patch
from unittest.mock import MagicMock, patch

from model_router import (
Limits,
Expand Down Expand Up @@ -39,6 +40,45 @@ def make_router(models=("cheap/small", "big/smart"), limits=None, **kw):


class ConfigTest(unittest.TestCase):
def test_timeout_reaches_both_routing_backends(self):
for backend, payload in [
("jev_api_key", {"code": 0, "data": {"decision": "cheap/small"}}),
("openrouter_api_key", {"answers": {"model": {"choice": "m0"}}}),
]:
for options, expected in [({}, 60), ({"timeout": 2.5}, 2.5)]:
with self.subTest(backend=backend, timeout=expected):
response = MagicMock()
response.__enter__.return_value = io.BytesIO(json.dumps(payload).encode())
with patch("urllib.request.urlopen", return_value=response) as urlopen:
router = make_router(**{backend: "test-key"}, **options)
self.assertEqual(router.route("hello"), "cheap/small")
self.assertEqual(urlopen.call_args.kwargs["timeout"], expected)

def test_timeout_reaches_catalog_download(self):
from model_router.catalog import refresh_catalog

refresh_catalog()
self.addCleanup(refresh_catalog)
response = MagicMock()
response.__enter__.return_value = io.BytesIO(
json.dumps(
{
"data": [
{
"id": "cheap/small",
"context_length": 8000,
"pricing": {"prompt": "0.0000001", "completion": "0.0000004"},
}
]
}
).encode()
)
with patch("urllib.request.urlopen", return_value=response) as urlopen:
router = Router(openrouter_api_key="test-key", models=["cheap/small"], timeout=5)
self.assertEqual(router.route("hello"), "cheap/small")
self.assertEqual(urlopen.call_count, 1)
self.assertEqual(urlopen.call_args.kwargs["timeout"], 5)

def test_needs_a_routing_key(self):
with self.assertRaises(RouterError):
make_router(openrouter_api_key=None)
Expand Down
Loading