Files
hermes-agent/tests/hermes_cli/test_model_cache_parallel_prefetch.py
ethernet 16652eea18 Merge remote-tracking branch 'origin/main' into ethie/pm-clean
# Conflicts:
#	gateway/config.py
#	gateway/config_loader.py
#	gateway/readiness.py
#	hermes_cli/managed_scope.py
#	hermes_cli/plugin_python_deps.py
#	hermes_cli/plugins_cmd.py
#	hermes_cli/update_cmd_maint.py
#	plugin-catalog/hindsight.yaml
#	plugins/plugin_loader.py
#	providers/__init__.py
#	scripts/run_tests.sh
#	tests/gateway/test_control_socket_windows_live.py
#	tests/gateway/test_gateway_streaming_nested_config.py
#	tests/hermes_cli/test_doctor.py
#	tests/hermes_cli/test_plan_reconciliation_windows_live.py
#	tests/hermes_cli/test_update_apply_shallow_count.py
#	tests/hermes_cli/test_update_concurrent_quarantine.py
#	tests/hermes_cli/test_update_shim_self_lock.py
#	tests/hermes_cli/test_verify_console_scripts.py
#	tests/tools/test_lazy_deps.py
#	tests/tui_gateway/test_subprocess_encoding.py
#	tools/lazy_deps.py
2026-09-23 15:26:34 -04:00

230 lines
10 KiB
Python

"""Tests for parallel model-catalog prefetch and thread-safe cache writes.
Regression tests for the serial /v1/models bottleneck: when the 1h disk cache
lapses, ``list_authenticated_providers()`` previously fetched each authed
provider's model list serially. With 10+ providers this stacked to 15-30s of
blocking HTTP round-trips. The parallel prefetch warms stale cache entries
concurrently via ThreadPoolExecutor before the serial picker loop starts.
"""
from __future__ import annotations
import time
from unittest.mock import patch
# ---------------------------------------------------------------------------
# Thread-safe cache entry update (hermes_cli/models.py)
# ---------------------------------------------------------------------------
class TestUpdateProviderCacheEntry:
"""Verify ``update_provider_cache_entry`` writes safely under concurrency."""
def test_writes_new_entry(self, tmp_path, monkeypatch):
"""A new entry is persisted to the cache file."""
import hermes_cli.models as mod
cache_path = tmp_path / "provider_models_cache.json"
monkeypatch.setattr(mod, "_provider_models_cache_path", lambda: cache_path)
with patch.object(mod, "_credential_fingerprint", return_value="fp1"):
mod.update_provider_cache_entry("openrouter", ["m1", "m2"])
cache = mod._load_provider_models_cache()
assert "openrouter" in cache
assert cache["openrouter"]["models"] == ["m1", "m2"]
assert cache["openrouter"]["fp"] == "fp1"
def test_does_not_clobber_other_entries(self, tmp_path, monkeypatch):
"""Concurrent writes to different providers don't lose entries."""
import hermes_cli.models as mod
cache_path = tmp_path / "provider_models_cache.json"
monkeypatch.setattr(mod, "_provider_models_cache_path", lambda: cache_path)
# Seed with one entry
with patch.object(mod, "_credential_fingerprint", return_value="fp_a"):
mod.update_provider_cache_entry("provider_a", ["a1"])
# Write a second entry
with patch.object(mod, "_credential_fingerprint", return_value="fp_b"):
mod.update_provider_cache_entry("provider_b", ["b1"])
cache = mod._load_provider_models_cache()
assert "provider_a" in cache
assert cache["provider_a"]["models"] == ["a1"]
assert "provider_b" in cache
assert cache["provider_b"]["models"] == ["b1"]
def test_skips_empty_models(self, tmp_path, monkeypatch):
"""Empty model lists are not written to cache."""
import hermes_cli.models as mod
cache_path = tmp_path / "provider_models_cache.json"
monkeypatch.setattr(mod, "_provider_models_cache_path", lambda: cache_path)
mod.update_provider_cache_entry("empty_provider", [])
cache = mod._load_provider_models_cache()
assert "empty_provider" not in cache
def test_concurrent_writes_no_lost_entries(self, tmp_path, monkeypatch):
"""Multiple threads writing different providers concurrently — all land."""
import hermes_cli.models as mod
import concurrent.futures
cache_path = tmp_path / "provider_models_cache.json"
monkeypatch.setattr(mod, "_provider_models_cache_path", lambda: cache_path)
providers = [f"prov_{i}" for i in range(10)]
with patch.object(mod, "_credential_fingerprint", side_effect=lambda p: f"fp_{p}"):
with concurrent.futures.ThreadPoolExecutor(max_workers=5) as executor:
list(executor.map(
lambda p: mod.update_provider_cache_entry(p, [f"model_{p}"]),
providers,
))
cache = mod._load_provider_models_cache()
for p in providers:
assert p in cache, f"{p} was lost in concurrent write"
assert cache[p]["models"] == [f"model_{p}"]
# ---------------------------------------------------------------------------
# Parallel prefetch (hermes_cli/model_switch.py)
# ---------------------------------------------------------------------------
class TestPrefetchProviderModelsParallel:
"""Verify ``_prefetch_provider_models_parallel`` fetches concurrently."""
def test_skips_all_fresh_entries(self, monkeypatch):
"""When all cache entries are fresh, no fetch is made."""
from hermes_cli.model_switch_providers import _prefetch_provider_models_parallel
fresh_cache = {
"openrouter": {"fp": "fp", "at": time.time(), "models": ["m1"]},
"anthropic": {"fp": "fp", "at": time.time(), "models": ["m2"]},
}
with patch("hermes_cli.models._load_provider_models_cache", return_value=fresh_cache), \
patch("hermes_cli.models._credential_fingerprint", return_value="fp"), \
patch("hermes_cli.models.cached_provider_model_ids") as fetch:
_prefetch_provider_models_parallel(["openrouter", "anthropic"])
fetch.assert_not_called()
def test_fetches_only_stale_entries(self, monkeypatch):
"""Only providers with stale/missing cache entries are fetched."""
from hermes_cli.model_switch_providers import _prefetch_provider_models_parallel
cache = {
"fresh_prov": {"fp": "fp_f", "at": time.time(), "models": ["m1"]},
}
fetch_calls = []
def mock_fetch(slug, force_refresh=False):
fetch_calls.append(slug)
return [f"model_{slug}"]
with patch("hermes_cli.models._load_provider_models_cache", return_value=cache), \
patch("hermes_cli.models._credential_fingerprint", return_value="fp_f"), \
patch("hermes_cli.models.cached_provider_model_ids", side_effect=mock_fetch), \
patch("hermes_cli.models.update_provider_cache_entry"):
_prefetch_provider_models_parallel(["fresh_prov", "stale_prov"])
assert "fresh_prov" not in fetch_calls
assert "stale_prov" in fetch_calls
def test_fetches_in_parallel(self, monkeypatch):
"""Multiple providers are fetched concurrently, not serially."""
from hermes_cli.model_switch_providers import _prefetch_provider_models_parallel
# Track overlap: if serial, no two fetches should overlap in time.
active = []
max_concurrent = [0]
lock = __import__("threading").Lock()
def mock_fetch(slug, force_refresh=False):
with lock:
active.append(slug)
max_concurrent[0] = max(max_concurrent[0], len(active))
time.sleep(0.05) # simulate network latency
with lock:
active.remove(slug)
return [f"model_{slug}"]
slugs = [f"prov_{i}" for i in range(6)]
with patch("hermes_cli.models._load_provider_models_cache", return_value={}), \
patch("hermes_cli.models._credential_fingerprint", return_value="fp"), \
patch("hermes_cli.models.cached_provider_model_ids", side_effect=mock_fetch), \
patch("hermes_cli.models.update_provider_cache_entry"):
_prefetch_provider_models_parallel(slugs)
assert max_concurrent[0] > 1, "fetches were serial, not parallel"
def test_swallows_exceptions(self):
"""A failing provider fetch doesn't raise — best-effort."""
from hermes_cli.model_switch_providers import _prefetch_provider_models_parallel
def mock_fetch(slug, force_refresh=False):
raise ConnectionError("simulated network failure")
with patch("hermes_cli.models._load_provider_models_cache", return_value={}), \
patch("hermes_cli.models._credential_fingerprint", return_value="fp"), \
patch("hermes_cli.models.cached_provider_model_ids", side_effect=mock_fetch), \
patch("hermes_cli.models.update_provider_cache_entry"):
# Should not raise
_prefetch_provider_models_parallel(["failing_prov"])
def test_empty_list_is_noop(self):
"""Empty provider list does nothing."""
from hermes_cli.model_switch_providers import _prefetch_provider_models_parallel
with patch("hermes_cli.models.cached_provider_model_ids") as fetch:
_prefetch_provider_models_parallel([])
fetch.assert_not_called()
def test_skips_ttl_expired_entries_the_serial_path_can_still_serve(self):
"""A TTL-expired entry inside the stale-serve window is not prefetched.
``cached_provider_model_ids`` returns such an entry from disk right
away and revalidates on a background thread, so blocking the picker
on a parallel fetch buys nothing. ``_PROVIDER_MODELS_STALE_SERVE_MAX``
is far longer than ``_PROVIDER_MODELS_CACHE_TTL``, so this is the
state every picker open a TTL after the previous one lands in.
"""
import hermes_cli.models as models_mod
from hermes_cli.model_switch_providers import _prefetch_provider_models_parallel
expired = time.time() - models_mod._PROVIDER_MODELS_CACHE_TTL - 60
cache = {"openrouter": {"fp": "fp", "at": expired, "models": ["m1"]}}
with patch("hermes_cli.models._load_provider_models_cache", return_value=cache), \
patch("hermes_cli.models._credential_fingerprint", return_value="fp"), \
patch("hermes_cli.models.cached_provider_model_ids") as fetch:
_prefetch_provider_models_parallel(["openrouter"])
fetch.assert_not_called()
def test_fetches_curated_fallback_rows_past_their_short_ttl(self):
"""A curated-fallback row is served only for ``_PROVIDER_MODELS_FALLBACK_TTL``
and never through the stale window, so the serial call blocks on it and the
parallel prefetch must fetch it."""
from hermes_cli.model_switch_providers import _prefetch_provider_models_parallel
cache = {"openrouter": {"fp": "fp", "at": time.time() - 7200, "models": ["m1"],
"fallback": True}}
fetched = []
def mock_fetch(slug, force_refresh=False):
fetched.append(slug)
return ["m1"]
with patch("hermes_cli.models._load_provider_models_cache", return_value=cache), \
patch("hermes_cli.models._credential_fingerprint", return_value="fp"), \
patch("hermes_cli.models.cached_provider_model_ids", side_effect=mock_fetch), \
patch("hermes_cli.models.update_provider_cache_entry"):
_prefetch_provider_models_parallel(["openrouter"])
assert fetched == ["openrouter"]