fix(gateway): keep models.dev refreshes off event loop
This commit is contained in:
@@ -118,7 +118,7 @@ def _provider_default_routes(provider: str) -> set[str]:
|
||||
from hermes_cli.providers import HERMES_OVERLAYS, get_provider
|
||||
|
||||
overlay = HERMES_OVERLAYS.get(provider)
|
||||
provider_def = get_provider(provider)
|
||||
provider_def = get_provider(provider, allow_network=False)
|
||||
for value in (
|
||||
getattr(overlay, "base_url_override", ""),
|
||||
getattr(provider_def, "base_url", ""),
|
||||
|
||||
@@ -20,6 +20,7 @@ rather than parsing the raw JSON themselves.
|
||||
|
||||
import json
|
||||
import logging
|
||||
import threading
|
||||
import time
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
@@ -33,10 +34,13 @@ logger = logging.getLogger(__name__)
|
||||
|
||||
MODELS_DEV_URL = "https://models.dev/api.json"
|
||||
_MODELS_DEV_CACHE_TTL = 3600 # 1 hour in-memory
|
||||
_MODELS_DEV_RETRY_DELAY = 300 # 5 minutes after a failed refresh
|
||||
|
||||
# In-memory cache
|
||||
_models_dev_cache: Dict[str, Any] = {}
|
||||
_models_dev_cache_time: float = 0
|
||||
_models_dev_retry_after: float = 0
|
||||
_models_dev_fetch_lock = threading.Lock()
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
@@ -237,7 +241,9 @@ def _save_disk_cache(data: Dict[str, Any]) -> None:
|
||||
logger.debug("Failed to save models.dev disk cache: %s", e)
|
||||
|
||||
|
||||
def fetch_models_dev(force_refresh: bool = False) -> Dict[str, Any]:
|
||||
def fetch_models_dev(
|
||||
force_refresh: bool = False, *, allow_network: bool = True
|
||||
) -> Dict[str, Any]:
|
||||
"""Fetch models.dev registry. Cache hierarchy: in-mem → disk → network.
|
||||
|
||||
Returns the full registry dict keyed by provider ID, or empty dict on failure.
|
||||
@@ -249,15 +255,28 @@ def fetch_models_dev(force_refresh: bool = False) -> Dict[str, Any]:
|
||||
``models.dev`` only changes when providers add new models, so a
|
||||
1 hour staleness window is acceptable (same TTL as in-mem cache).
|
||||
3. Network fetch → on success, save to disk + in-mem and return.
|
||||
4. Network fails → fall back to ANY available disk cache (even stale)
|
||||
with a short 5 min in-mem grace period before retrying network.
|
||||
4. Network fails → fall back to ANY available cache (even stale) and
|
||||
suppress additional automatic refreshes for 5 minutes.
|
||||
|
||||
When ``force_refresh=True`` (used by ``hermes config refresh``, the
|
||||
\"refresh model catalog\" code path), stages 1 and 2 are skipped. The
|
||||
function always hits the network and only falls back to disk if the
|
||||
network call fails.
|
||||
function bypasses caches and failure backoff, then only falls back to
|
||||
cached data if the network call fails. When ``allow_network=False``, any
|
||||
memory or disk cache is returned regardless of age and no request is made.
|
||||
"""
|
||||
global _models_dev_cache, _models_dev_cache_time
|
||||
global _models_dev_cache, _models_dev_cache_time, _models_dev_retry_after
|
||||
|
||||
if not allow_network:
|
||||
if _models_dev_cache:
|
||||
return _models_dev_cache
|
||||
disk_data = _load_disk_cache()
|
||||
if disk_data:
|
||||
_models_dev_cache = disk_data
|
||||
disk_age = _disk_cache_age_seconds()
|
||||
_models_dev_cache_time = (
|
||||
time.time() - disk_age if disk_age is not None else 0
|
||||
)
|
||||
return _models_dev_cache
|
||||
|
||||
# Stage 1: fresh in-memory cache wins. This is the hot path on
|
||||
# long-lived processes — no I/O, no system calls.
|
||||
@@ -288,34 +307,72 @@ def fetch_models_dev(force_refresh: bool = False) -> Dict[str, Any]:
|
||||
)
|
||||
return _models_dev_cache
|
||||
|
||||
# Stage 3: network fetch.
|
||||
try:
|
||||
response = requests.get(MODELS_DEV_URL, timeout=15)
|
||||
response.raise_for_status()
|
||||
data = response.json()
|
||||
if isinstance(data, dict) and data:
|
||||
# Failed automatic refreshes are process-wide. Avoid making every caller
|
||||
# retry the same unreachable endpoint while stale data remains usable.
|
||||
if not force_refresh and time.time() < _models_dev_retry_after:
|
||||
if not _models_dev_cache:
|
||||
_models_dev_cache = _load_disk_cache()
|
||||
_models_dev_cache_time = 0
|
||||
return _models_dev_cache
|
||||
|
||||
# Stage 3: singleflight network fetch. Recheck state after acquiring the
|
||||
# lock because another caller may have refreshed or established backoff.
|
||||
with _models_dev_fetch_lock:
|
||||
now = time.time()
|
||||
if not force_refresh:
|
||||
if (
|
||||
_models_dev_cache
|
||||
and (now - _models_dev_cache_time) < _MODELS_DEV_CACHE_TTL
|
||||
):
|
||||
return _models_dev_cache
|
||||
if now < _models_dev_retry_after:
|
||||
if not _models_dev_cache:
|
||||
_models_dev_cache = _load_disk_cache()
|
||||
_models_dev_cache_time = 0
|
||||
return _models_dev_cache
|
||||
|
||||
try:
|
||||
response = requests.get(MODELS_DEV_URL, timeout=15)
|
||||
response.raise_for_status()
|
||||
data = response.json()
|
||||
if not isinstance(data, dict) or not data:
|
||||
raise ValueError("models.dev returned an empty or invalid registry")
|
||||
|
||||
_save_disk_cache(data)
|
||||
_models_dev_cache = data
|
||||
_models_dev_cache_time = time.time()
|
||||
_save_disk_cache(data)
|
||||
_models_dev_retry_after = 0
|
||||
logger.debug(
|
||||
"Fetched models.dev registry: %d providers, %d total models",
|
||||
len(data),
|
||||
sum(len(p.get("models", {})) for p in data.values() if isinstance(p, dict)),
|
||||
sum(
|
||||
len(p.get("models", {}))
|
||||
for p in data.values()
|
||||
if isinstance(p, dict)
|
||||
),
|
||||
)
|
||||
return data
|
||||
except Exception as e:
|
||||
logger.debug("Failed to fetch models.dev: %s", e)
|
||||
except Exception as e:
|
||||
_models_dev_retry_after = time.time() + _MODELS_DEV_RETRY_DELAY
|
||||
logger.debug(
|
||||
"Failed to fetch models.dev; retry suppressed for %ds: %s",
|
||||
_MODELS_DEV_RETRY_DELAY,
|
||||
e,
|
||||
)
|
||||
|
||||
# Stage 4: network failed — fall back to whatever disk cache exists,
|
||||
# even if it's stale. Give it a short 5 min in-mem TTL so we retry
|
||||
# the network soon instead of serving stale data for a full hour.
|
||||
if not _models_dev_cache:
|
||||
_models_dev_cache = _load_disk_cache()
|
||||
if _models_dev_cache:
|
||||
_models_dev_cache_time = time.time() - _MODELS_DEV_CACHE_TTL + 300
|
||||
logger.debug("Loaded models.dev from disk cache (%d providers)", len(_models_dev_cache))
|
||||
# Stage 4: network failed — return any stale memory/disk cache. Cache
|
||||
# freshness remains expired; the retry-after timestamp controls when
|
||||
# the next automatic request is allowed.
|
||||
if not _models_dev_cache:
|
||||
_models_dev_cache = _load_disk_cache()
|
||||
_models_dev_cache_time = 0
|
||||
if _models_dev_cache:
|
||||
logger.debug(
|
||||
"Loaded stale models.dev disk cache (%d providers)",
|
||||
len(_models_dev_cache),
|
||||
)
|
||||
|
||||
return _models_dev_cache
|
||||
return _models_dev_cache
|
||||
|
||||
|
||||
def lookup_models_dev_context(provider: str, model: str) -> Optional[int]:
|
||||
@@ -671,7 +728,9 @@ def _parse_provider_info(provider_id: str, raw: Dict[str, Any]) -> ProviderInfo:
|
||||
# Provider-level queries
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
def get_provider_info(provider_id: str) -> Optional[ProviderInfo]:
|
||||
def get_provider_info(
|
||||
provider_id: str, *, allow_network: bool = True
|
||||
) -> Optional[ProviderInfo]:
|
||||
"""Get full provider metadata from models.dev.
|
||||
|
||||
Accepts either a Hermes provider ID (e.g. "kilocode") or a models.dev
|
||||
@@ -680,7 +739,11 @@ def get_provider_info(provider_id: str) -> Optional[ProviderInfo]:
|
||||
# Resolve Hermes ID → models.dev ID
|
||||
mdev_id = PROVIDER_TO_MODELS_DEV.get(provider_id, provider_id)
|
||||
|
||||
data = fetch_models_dev()
|
||||
data = (
|
||||
fetch_models_dev()
|
||||
if allow_network
|
||||
else fetch_models_dev(allow_network=False)
|
||||
)
|
||||
raw = data.get(mdev_id)
|
||||
if not isinstance(raw, dict):
|
||||
return None
|
||||
|
||||
@@ -13131,7 +13131,8 @@ class GatewayRunner(GatewayAuthorizationMixin, GatewayKanbanWatchersMixin, Gatew
|
||||
try:
|
||||
from hermes_cli.route_identity import should_clear_context_pin
|
||||
|
||||
if should_clear_context_pin(
|
||||
if await asyncio.to_thread(
|
||||
should_clear_context_pin,
|
||||
None, # model match already checked above
|
||||
None,
|
||||
_msg_model_cfg.get("base_url"),
|
||||
@@ -13726,7 +13727,8 @@ class GatewayRunner(GatewayAuthorizationMixin, GatewayKanbanWatchersMixin, Gatew
|
||||
try:
|
||||
from hermes_cli.route_identity import should_clear_context_pin
|
||||
|
||||
if should_clear_context_pin(
|
||||
if await asyncio.to_thread(
|
||||
should_clear_context_pin,
|
||||
_hyg_configured_model,
|
||||
_hyg_model,
|
||||
_hyg_configured_base_url,
|
||||
|
||||
@@ -2010,7 +2010,8 @@ class GatewaySlashCommandsMixin:
|
||||
try:
|
||||
from hermes_cli.route_identity import should_clear_context_pin
|
||||
|
||||
if should_clear_context_pin(
|
||||
if await asyncio.to_thread(
|
||||
should_clear_context_pin,
|
||||
_persist_model_cfg.get("default")
|
||||
or _persist_model_cfg.get("model"),
|
||||
result.new_model,
|
||||
@@ -2338,7 +2339,8 @@ class GatewaySlashCommandsMixin:
|
||||
try:
|
||||
from hermes_cli.route_identity import should_clear_context_pin
|
||||
|
||||
if should_clear_context_pin(
|
||||
if await asyncio.to_thread(
|
||||
should_clear_context_pin,
|
||||
model_cfg.get("default") or model_cfg.get("model"),
|
||||
result.new_model,
|
||||
model_cfg.get("base_url"),
|
||||
|
||||
@@ -431,7 +431,7 @@ def normalize_provider(name: str) -> str:
|
||||
return ALIASES.get(key, key)
|
||||
|
||||
|
||||
def get_provider(name: str) -> Optional[ProviderDef]:
|
||||
def get_provider(name: str, *, allow_network: bool = True) -> Optional[ProviderDef]:
|
||||
"""Look up a built-in provider by id or alias.
|
||||
|
||||
Resolution order:
|
||||
@@ -450,7 +450,11 @@ def get_provider(name: str) -> Optional[ProviderDef]:
|
||||
# Try to get models.dev data
|
||||
try:
|
||||
from agent.models_dev import get_provider_info as _mdev_provider
|
||||
mdev_info = _mdev_provider(canonical)
|
||||
mdev_info = (
|
||||
_mdev_provider(canonical)
|
||||
if allow_network
|
||||
else _mdev_provider(canonical, allow_network=False)
|
||||
)
|
||||
except Exception:
|
||||
mdev_info = None
|
||||
|
||||
|
||||
@@ -1,11 +1,17 @@
|
||||
"""Tests for agent.models_dev — models.dev registry integration."""
|
||||
import threading
|
||||
import time
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
from unittest.mock import patch, MagicMock
|
||||
|
||||
import pytest
|
||||
|
||||
from agent.models_dev import (
|
||||
PROVIDER_TO_MODELS_DEV,
|
||||
_extract_context,
|
||||
fetch_models_dev,
|
||||
get_model_capabilities,
|
||||
get_provider_info,
|
||||
lookup_models_dev_context,
|
||||
)
|
||||
|
||||
@@ -164,6 +170,18 @@ class TestLookupModelsDevContext:
|
||||
|
||||
|
||||
class TestFetchModelsDev:
|
||||
@pytest.fixture(autouse=True)
|
||||
def _reset_fetch_state(self):
|
||||
import agent.models_dev as md
|
||||
|
||||
md._models_dev_cache = {}
|
||||
md._models_dev_cache_time = 0
|
||||
md._models_dev_retry_after = 0
|
||||
yield
|
||||
md._models_dev_cache = {}
|
||||
md._models_dev_cache_time = 0
|
||||
md._models_dev_retry_after = 0
|
||||
|
||||
@patch("agent.models_dev.requests.get")
|
||||
def test_fetch_success(self, mock_get):
|
||||
mock_resp = MagicMock()
|
||||
@@ -303,6 +321,158 @@ class TestFetchModelsDev:
|
||||
mock_get.assert_called_once()
|
||||
assert "anthropic" in result
|
||||
|
||||
@patch("agent.models_dev.requests.get")
|
||||
def test_stale_cache_failure_enters_backoff_and_suppresses_retry(self, mock_get):
|
||||
import agent.models_dev as md
|
||||
|
||||
mock_get.side_effect = OSError("models.dev unreachable")
|
||||
md._models_dev_cache = SAMPLE_REGISTRY
|
||||
md._models_dev_cache_time = time.time() - md._MODELS_DEV_CACHE_TTL - 1
|
||||
|
||||
with patch.object(
|
||||
md,
|
||||
"_disk_cache_age_seconds",
|
||||
return_value=md._MODELS_DEV_CACHE_TTL + 60,
|
||||
), patch.object(md, "_load_disk_cache", return_value=SAMPLE_REGISTRY):
|
||||
first = fetch_models_dev()
|
||||
second = fetch_models_dev()
|
||||
|
||||
assert first == SAMPLE_REGISTRY
|
||||
assert second == SAMPLE_REGISTRY
|
||||
assert md._models_dev_retry_after > time.time()
|
||||
mock_get.assert_called_once()
|
||||
|
||||
@patch("agent.models_dev.requests.get")
|
||||
def test_missing_cache_failure_enters_backoff(self, mock_get):
|
||||
import agent.models_dev as md
|
||||
|
||||
mock_get.side_effect = OSError("models.dev unreachable")
|
||||
with patch.object(md, "_disk_cache_age_seconds", return_value=None), patch.object(
|
||||
md, "_load_disk_cache", return_value={}
|
||||
):
|
||||
first = fetch_models_dev()
|
||||
second = fetch_models_dev()
|
||||
|
||||
assert first == {}
|
||||
assert second == {}
|
||||
assert md._models_dev_retry_after > time.time()
|
||||
mock_get.assert_called_once()
|
||||
|
||||
@patch("agent.models_dev.requests.get")
|
||||
def test_concurrent_refreshes_share_one_network_request(self, mock_get):
|
||||
import agent.models_dev as md
|
||||
|
||||
request_started = threading.Event()
|
||||
release_request = threading.Event()
|
||||
response = MagicMock()
|
||||
response.json.return_value = SAMPLE_REGISTRY
|
||||
|
||||
def blocking_get(*_args, **_kwargs):
|
||||
request_started.set()
|
||||
assert release_request.wait(timeout=5)
|
||||
return response
|
||||
|
||||
mock_get.side_effect = blocking_get
|
||||
with patch.object(md, "_disk_cache_age_seconds", return_value=None), patch.object(
|
||||
md, "_save_disk_cache"
|
||||
), ThreadPoolExecutor(max_workers=6) as pool:
|
||||
futures = [pool.submit(fetch_models_dev) for _ in range(6)]
|
||||
assert request_started.wait(timeout=2)
|
||||
release_request.set()
|
||||
results = [future.result(timeout=5) for future in futures]
|
||||
|
||||
assert results == [SAMPLE_REGISTRY] * 6
|
||||
mock_get.assert_called_once()
|
||||
|
||||
@patch("agent.models_dev.requests.get")
|
||||
def test_force_refresh_bypasses_failure_backoff(self, mock_get):
|
||||
import agent.models_dev as md
|
||||
|
||||
response = MagicMock()
|
||||
response.json.return_value = SAMPLE_REGISTRY
|
||||
mock_get.side_effect = [OSError("models.dev unreachable"), response]
|
||||
|
||||
with patch.object(md, "_disk_cache_age_seconds", return_value=None), patch.object(
|
||||
md, "_load_disk_cache", return_value={}
|
||||
), patch.object(md, "_save_disk_cache"):
|
||||
assert fetch_models_dev() == {}
|
||||
assert fetch_models_dev(force_refresh=True) == SAMPLE_REGISTRY
|
||||
|
||||
assert mock_get.call_count == 2
|
||||
assert md._models_dev_retry_after == 0
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("cache", "cache_time", "disk_data", "expected"),
|
||||
[
|
||||
(SAMPLE_REGISTRY, lambda md: time.time(), {}, SAMPLE_REGISTRY),
|
||||
(
|
||||
SAMPLE_REGISTRY,
|
||||
lambda md: time.time() - md._MODELS_DEV_CACHE_TTL - 1,
|
||||
{},
|
||||
SAMPLE_REGISTRY,
|
||||
),
|
||||
({}, lambda _md: 0, {}, {}),
|
||||
],
|
||||
ids=["fresh-memory", "stale-memory", "missing"],
|
||||
)
|
||||
@patch("agent.models_dev.requests.get")
|
||||
def test_network_disabled_never_fetches(
|
||||
self, mock_get, cache, cache_time, disk_data, expected
|
||||
):
|
||||
import agent.models_dev as md
|
||||
|
||||
md._models_dev_cache = cache
|
||||
md._models_dev_cache_time = cache_time(md)
|
||||
with patch.object(md, "_load_disk_cache", return_value=disk_data):
|
||||
result = fetch_models_dev(allow_network=False)
|
||||
|
||||
assert result == expected
|
||||
mock_get.assert_not_called()
|
||||
|
||||
@patch("agent.models_dev.requests.get")
|
||||
def test_network_disabled_loads_stale_disk_cache(self, mock_get):
|
||||
import agent.models_dev as md
|
||||
|
||||
with patch.object(md, "_load_disk_cache", return_value=SAMPLE_REGISTRY):
|
||||
result = fetch_models_dev(allow_network=False)
|
||||
|
||||
assert result == SAMPLE_REGISTRY
|
||||
mock_get.assert_not_called()
|
||||
|
||||
@patch("agent.models_dev.fetch_models_dev", return_value=SAMPLE_REGISTRY)
|
||||
def test_provider_info_propagates_network_disabled(self, mock_fetch):
|
||||
info = get_provider_info("anthropic", allow_network=False)
|
||||
|
||||
assert info is not None
|
||||
mock_fetch.assert_called_once_with(allow_network=False)
|
||||
|
||||
@patch("agent.models_dev.fetch_models_dev", return_value=SAMPLE_REGISTRY)
|
||||
def test_provider_info_default_preserves_zero_argument_fetch(self, mock_fetch):
|
||||
info = get_provider_info("anthropic")
|
||||
|
||||
assert info is not None
|
||||
mock_fetch.assert_called_once_with()
|
||||
|
||||
def test_provider_definition_propagates_network_disabled(self):
|
||||
from hermes_cli.providers import get_provider
|
||||
|
||||
with patch(
|
||||
"agent.models_dev.get_provider_info", return_value=None
|
||||
) as mock_provider_info:
|
||||
get_provider("anthropic", allow_network=False)
|
||||
|
||||
mock_provider_info.assert_called_once_with(
|
||||
"anthropic", allow_network=False
|
||||
)
|
||||
|
||||
def test_default_route_lookup_is_cache_only(self):
|
||||
from agent.agent_init import _provider_default_routes
|
||||
|
||||
with patch("hermes_cli.providers.get_provider", return_value=None) as mock_get:
|
||||
_provider_default_routes("anthropic")
|
||||
|
||||
mock_get.assert_called_once_with("anthropic", allow_network=False)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# get_model_capabilities — vision via modalities.input
|
||||
|
||||
@@ -223,3 +223,91 @@ async def test_prepare_image_routing_runs_off_the_event_loop(monkeypatch):
|
||||
"the blocking image-routing decision must be offloaded off the gateway "
|
||||
"event loop, not run inline on it"
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_prepare_route_identity_check_keeps_event_loop_responsive(monkeypatch):
|
||||
"""A slow route-identity check must not block gateway heartbeats."""
|
||||
import asyncio
|
||||
import threading
|
||||
from types import SimpleNamespace
|
||||
|
||||
runner = _make_runner()
|
||||
source = _source()
|
||||
event = MessageEvent(
|
||||
text="inspect @AGENTS.md",
|
||||
message_type=MessageType.TEXT,
|
||||
source=source,
|
||||
)
|
||||
started = threading.Event()
|
||||
released_by_event_loop = threading.Event()
|
||||
seen = {}
|
||||
main_thread = threading.current_thread()
|
||||
|
||||
cfg = {
|
||||
"model": {
|
||||
"default": "test-model",
|
||||
"provider": "test-provider",
|
||||
"base_url": "https://example.invalid/v1",
|
||||
"context_length": 128000,
|
||||
}
|
||||
}
|
||||
monkeypatch.setattr("gateway.run._load_gateway_config", lambda: cfg)
|
||||
monkeypatch.setattr(
|
||||
runner,
|
||||
"_resolve_session_agent_runtime",
|
||||
lambda **_kwargs: (
|
||||
"test-model",
|
||||
{
|
||||
"provider": "test-provider",
|
||||
"base_url": "https://example.invalid/v1",
|
||||
"api_key": "",
|
||||
},
|
||||
),
|
||||
)
|
||||
|
||||
def blocking_route_identity_check(*_args):
|
||||
seen["thread"] = threading.current_thread()
|
||||
started.set()
|
||||
seen["event_loop_progressed"] = released_by_event_loop.wait(timeout=2)
|
||||
return False
|
||||
|
||||
monkeypatch.setattr(
|
||||
"hermes_cli.route_identity.should_clear_context_pin",
|
||||
blocking_route_identity_check,
|
||||
)
|
||||
|
||||
async def fake_context_length(*_args, **_kwargs):
|
||||
return 128000
|
||||
|
||||
async def fake_preprocess(message, **_kwargs):
|
||||
return SimpleNamespace(
|
||||
blocked=False,
|
||||
expanded=False,
|
||||
message=message,
|
||||
warnings=[],
|
||||
)
|
||||
|
||||
monkeypatch.setattr(
|
||||
"agent.model_metadata.get_model_context_length_async", fake_context_length
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
"agent.context_references.preprocess_context_references_async",
|
||||
fake_preprocess,
|
||||
)
|
||||
|
||||
async def heartbeat_ticker():
|
||||
while not started.is_set():
|
||||
await asyncio.sleep(0)
|
||||
await asyncio.sleep(0)
|
||||
released_by_event_loop.set()
|
||||
|
||||
heartbeat = asyncio.create_task(heartbeat_ticker())
|
||||
result = await runner._prepare_inbound_message_text(
|
||||
event=event, source=source, history=[]
|
||||
)
|
||||
await heartbeat
|
||||
|
||||
assert result == "inspect @AGENTS.md"
|
||||
assert seen["event_loop_progressed"] is True
|
||||
assert seen["thread"] is not main_thread
|
||||
|
||||
Reference in New Issue
Block a user