diff --git a/agent/agent_init.py b/agent/agent_init.py index 0fafd50bb5..ea473632c6 100644 --- a/agent/agent_init.py +++ b/agent/agent_init.py @@ -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", ""), diff --git a/agent/models_dev.py b/agent/models_dev.py index 590f77806a..2abfef675b 100644 --- a/agent/models_dev.py +++ b/agent/models_dev.py @@ -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 diff --git a/gateway/run.py b/gateway/run.py index 78ca83129e..3fbcb43cf8 100644 --- a/gateway/run.py +++ b/gateway/run.py @@ -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, diff --git a/gateway/slash_commands.py b/gateway/slash_commands.py index 90a954f7aa..43daaaa3c9 100644 --- a/gateway/slash_commands.py +++ b/gateway/slash_commands.py @@ -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"), diff --git a/hermes_cli/providers.py b/hermes_cli/providers.py index 5515a16700..53531820b7 100644 --- a/hermes_cli/providers.py +++ b/hermes_cli/providers.py @@ -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 diff --git a/tests/agent/test_models_dev.py b/tests/agent/test_models_dev.py index b4bbbf753d..56592f0bcd 100644 --- a/tests/agent/test_models_dev.py +++ b/tests/agent/test_models_dev.py @@ -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 diff --git a/tests/gateway/test_image_input_routing_runtime.py b/tests/gateway/test_image_input_routing_runtime.py index 8ce68aeaac..2585430d0e 100644 --- a/tests/gateway/test_image_input_routing_runtime.py +++ b/tests/gateway/test_image_input_routing_runtime.py @@ -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