"""Tests for the inference registry resolver. Validates: - Resolution returns correct target from mocked DB records - TTL expiry triggers re-resolution - Invalidation clears cached entries - Missing binding raises typed error (fail-closed) - Disabled endpoint raises typed error - auth_secret_ref is preserved as-is (not resolved during caching) - Deterministic: same input always returns same output Requirements: 3.5, 3.9 """ from __future__ import annotations import time from typing import Any from unittest.mock import patch from uuid import UUID, uuid4 import pytest from services.shared.inference.errors import InferenceError, InferenceErrorCategory from services.shared.inference.models import InferenceTarget from services.shared.inference.registry import ( RegistryCache, RegistryDB, RegistryResolver, ) # --------------------------------------------------------------------------- # Mock DB implementation # --------------------------------------------------------------------------- class MockRegistryDB(RegistryDB): """In-memory mock of the registry database for testing.""" def __init__(self) -> None: self.bindings: dict[tuple[UUID, str], dict[str, Any]] = {} self.deployments: dict[UUID, dict[str, Any]] = {} self.endpoints: dict[UUID, dict[str, Any]] = {} self.call_count: dict[str, int] = { "get_active_binding": 0, "get_model_deployment": 0, "get_inference_endpoint": 0, } async def get_active_binding( self, agent_id: UUID, stage: str ) -> dict[str, Any] | None: self.call_count["get_active_binding"] += 1 return self.bindings.get((agent_id, stage)) async def get_model_deployment( self, deployment_id: UUID ) -> dict[str, Any] | None: self.call_count["get_model_deployment"] += 1 return self.deployments.get(deployment_id) async def get_inference_endpoint( self, endpoint_id: UUID ) -> dict[str, Any] | None: self.call_count["get_inference_endpoint"] += 1 return self.endpoints.get(endpoint_id) # --------------------------------------------------------------------------- # Fixtures # --------------------------------------------------------------------------- def _build_test_db() -> tuple[MockRegistryDB, UUID, UUID, UUID, UUID]: """Build a mock DB with a complete resolution chain. Returns (db, agent_id, endpoint_id, deployment_id, binding_id). """ db = MockRegistryDB() agent_id = uuid4() endpoint_id = uuid4() deployment_id = uuid4() binding_id = uuid4() db.endpoints[endpoint_id] = { "id": endpoint_id, "name": "vllm-adjudicator", "protocol": "openai_chat", "base_url": "http://vllm.stonks-oracle.svc:8000", "auth_secret_ref": "VLLM_API_KEY", "auth_scheme": "bearer", "default_headers": {"X-Request-Source": "stonks-oracle"}, "health_path": "/health", "enabled": True, "revision": 1, } db.deployments[deployment_id] = { "id": deployment_id, "endpoint_id": endpoint_id, "served_model_name": "stonks-adjudicator-9b", "display_name": "Qwen3.5-9B Adjudicator", "capabilities": { "chat_completions": True, "json_schema": True, "usage": True, "seed": True, }, "context_window": 8192, "max_output_tokens": 1536, "quantization": "NVFP4", "runtime_metadata": {"extra_body": {"guided_decoding_backend": "outlines"}}, "enabled": True, "revision": 1, } db.bindings[(agent_id, "extraction")] = { "id": binding_id, "agent_id": agent_id, "stage": "extraction", "model_deployment_id": deployment_id, "route_order": 0, "routing_config": {}, "is_active": True, "revision": 1, } return db, agent_id, endpoint_id, deployment_id, binding_id # --------------------------------------------------------------------------- # RegistryCache tests # --------------------------------------------------------------------------- class TestRegistryCache: """Tests for the TTL cache implementation.""" def test_set_and_get(self) -> None: """Basic set/get returns stored value.""" cache = RegistryCache(ttl_seconds=60.0) cache.set("binding:abc:extraction", {"target": "value"}) assert cache.get("binding:abc:extraction") == {"target": "value"} def test_get_missing_key_returns_none(self) -> None: """Missing key returns None.""" cache = RegistryCache(ttl_seconds=60.0) assert cache.get("nonexistent") is None def test_ttl_expiry(self) -> None: """Expired entries return None.""" cache = RegistryCache(ttl_seconds=0.01) # 10ms TTL cache.set("key", "value") time.sleep(0.02) # Wait for expiry assert cache.get("key") is None def test_invalidate_exact_key(self) -> None: """Invalidate removes exact matching key.""" cache = RegistryCache(ttl_seconds=60.0) cache.set("endpoint:abc-123", {"data": 1}) cache.set("endpoint:def-456", {"data": 2}) cache.invalidate("endpoint:abc-123") assert cache.get("endpoint:abc-123") is None assert cache.get("endpoint:def-456") == {"data": 2} def test_invalidate_prefix(self) -> None: """Invalidate with prefix removes all matching entries.""" cache = RegistryCache(ttl_seconds=60.0) cache.set("binding:agent1:extraction", "t1") cache.set("binding:agent1:sentiment", "t2") cache.set("binding:agent2:extraction", "t3") cache.invalidate("binding:agent1:") assert cache.get("binding:agent1:extraction") is None assert cache.get("binding:agent1:sentiment") is None assert cache.get("binding:agent2:extraction") == "t3" def test_clear_removes_all(self) -> None: """Clear removes all entries.""" cache = RegistryCache(ttl_seconds=60.0) cache.set("a", 1) cache.set("b", 2) cache.clear() assert len(cache) == 0 assert cache.get("a") is None assert cache.get("b") is None def test_contains_operator(self) -> None: """__contains__ checks non-expired existence.""" cache = RegistryCache(ttl_seconds=60.0) cache.set("present", "yes") assert "present" in cache assert "absent" not in cache # --------------------------------------------------------------------------- # RegistryResolver tests — happy path # --------------------------------------------------------------------------- class TestResolverHappyPath: """Resolution returns correct target from mocked DB records.""" @pytest.mark.asyncio async def test_resolve_returns_correct_target(self) -> None: """Full resolution chain produces correct InferenceTarget.""" db, agent_id, endpoint_id, deployment_id, _ = _build_test_db() resolver = RegistryResolver(db, cache_ttl_seconds=60.0) target = await resolver.resolve_target(agent_id, "extraction") assert isinstance(target, InferenceTarget) assert target.endpoint_id == endpoint_id assert target.deployment_id == deployment_id assert target.protocol == "openai_chat" assert target.base_url == "http://vllm.stonks-oracle.svc:8000" assert target.model == "stonks-adjudicator-9b" assert target.capabilities.chat_completions is True assert target.capabilities.json_schema is True assert target.capabilities.usage is True assert target.capabilities.seed is True assert target.capabilities.json_object is False assert target.context_window == 8192 assert target.max_output_tokens == 1536 assert target.extra_headers == {"X-Request-Source": "stonks-oracle"} assert target.extra_body == {"guided_decoding_backend": "outlines"} @pytest.mark.asyncio async def test_resolve_caches_target(self) -> None: """Second resolution uses cache instead of querying DB.""" db, agent_id, *_ = _build_test_db() resolver = RegistryResolver(db, cache_ttl_seconds=60.0) # First call queries DB await resolver.resolve_target(agent_id, "extraction") assert db.call_count["get_active_binding"] == 1 # Second call uses cache await resolver.resolve_target(agent_id, "extraction") assert db.call_count["get_active_binding"] == 1 # Not incremented # --------------------------------------------------------------------------- # RegistryResolver tests — TTL expiry # --------------------------------------------------------------------------- class TestResolverTTLExpiry: """TTL expiry triggers re-resolution.""" @pytest.mark.asyncio async def test_expired_cache_re_resolves(self) -> None: """After TTL expires, resolver queries DB again.""" db, agent_id, *_ = _build_test_db() resolver = RegistryResolver(db, cache_ttl_seconds=0.01) # First resolution await resolver.resolve_target(agent_id, "extraction") assert db.call_count["get_active_binding"] == 1 # Wait for TTL to expire time.sleep(0.02) # Should re-query await resolver.resolve_target(agent_id, "extraction") assert db.call_count["get_active_binding"] == 2 # --------------------------------------------------------------------------- # RegistryResolver tests — invalidation # --------------------------------------------------------------------------- class TestResolverInvalidation: """Invalidation clears cached entries.""" @pytest.mark.asyncio async def test_invalidate_endpoint_clears_cache(self) -> None: """invalidate(endpoint_id) forces re-resolution.""" db, agent_id, endpoint_id, *_ = _build_test_db() resolver = RegistryResolver(db, cache_ttl_seconds=60.0) # Populate cache await resolver.resolve_target(agent_id, "extraction") assert db.call_count["get_active_binding"] == 1 # Invalidate resolver.invalidate(endpoint_id) # Should re-query await resolver.resolve_target(agent_id, "extraction") assert db.call_count["get_active_binding"] == 2 @pytest.mark.asyncio async def test_invalidate_all_clears_everything(self) -> None: """invalidate_all() clears all cache entries.""" db, agent_id, *_ = _build_test_db() resolver = RegistryResolver(db, cache_ttl_seconds=60.0) # Populate cache await resolver.resolve_target(agent_id, "extraction") # Full clear resolver.invalidate_all() # Should re-query await resolver.resolve_target(agent_id, "extraction") assert db.call_count["get_active_binding"] == 2 # --------------------------------------------------------------------------- # RegistryResolver tests — fail-closed behavior # --------------------------------------------------------------------------- class TestResolverFailClosed: """Missing or disabled resources raise typed errors.""" @pytest.mark.asyncio async def test_missing_binding_raises_capability_unavailable(self) -> None: """No active binding raises InferenceError.""" db = MockRegistryDB() resolver = RegistryResolver(db, cache_ttl_seconds=60.0) with pytest.raises(InferenceError) as exc_info: await resolver.resolve_target(uuid4(), "extraction") assert exc_info.value.category == InferenceErrorCategory.CAPABILITY_UNAVAILABLE assert "No active binding" in str(exc_info.value) @pytest.mark.asyncio async def test_inactive_binding_raises_capability_unavailable(self) -> None: """Inactive binding raises InferenceError.""" db, agent_id, *_ = _build_test_db() # Mark binding inactive db.bindings[(agent_id, "extraction")]["is_active"] = False resolver = RegistryResolver(db, cache_ttl_seconds=60.0) with pytest.raises(InferenceError) as exc_info: await resolver.resolve_target(agent_id, "extraction") assert exc_info.value.category == InferenceErrorCategory.CAPABILITY_UNAVAILABLE @pytest.mark.asyncio async def test_missing_deployment_raises_capability_unavailable(self) -> None: """Missing model deployment raises InferenceError.""" db, agent_id, *_ = _build_test_db() # Remove the deployment db.deployments.clear() resolver = RegistryResolver(db, cache_ttl_seconds=60.0) with pytest.raises(InferenceError) as exc_info: await resolver.resolve_target(agent_id, "extraction") assert exc_info.value.category == InferenceErrorCategory.CAPABILITY_UNAVAILABLE assert "not found" in str(exc_info.value) @pytest.mark.asyncio async def test_disabled_deployment_raises_capability_unavailable(self) -> None: """Disabled model deployment raises InferenceError.""" db, agent_id, _, deployment_id, _ = _build_test_db() db.deployments[deployment_id]["enabled"] = False resolver = RegistryResolver(db, cache_ttl_seconds=60.0) with pytest.raises(InferenceError) as exc_info: await resolver.resolve_target(agent_id, "extraction") assert exc_info.value.category == InferenceErrorCategory.CAPABILITY_UNAVAILABLE assert "disabled" in str(exc_info.value) @pytest.mark.asyncio async def test_missing_endpoint_raises_capability_unavailable(self) -> None: """Missing inference endpoint raises InferenceError.""" db, agent_id, *_ = _build_test_db() # Remove the endpoint db.endpoints.clear() resolver = RegistryResolver(db, cache_ttl_seconds=60.0) with pytest.raises(InferenceError) as exc_info: await resolver.resolve_target(agent_id, "extraction") assert exc_info.value.category == InferenceErrorCategory.CAPABILITY_UNAVAILABLE assert "not found" in str(exc_info.value) @pytest.mark.asyncio async def test_disabled_endpoint_raises_capability_unavailable(self) -> None: """Disabled inference endpoint raises InferenceError.""" db, agent_id, endpoint_id, *_ = _build_test_db() db.endpoints[endpoint_id]["enabled"] = False resolver = RegistryResolver(db, cache_ttl_seconds=60.0) with pytest.raises(InferenceError) as exc_info: await resolver.resolve_target(agent_id, "extraction") assert exc_info.value.category == InferenceErrorCategory.CAPABILITY_UNAVAILABLE assert "disabled" in str(exc_info.value) @pytest.mark.asyncio async def test_binding_with_no_deployment_id_raises(self) -> None: """Binding with model_deployment_id=None raises InferenceError.""" db, agent_id, *_ = _build_test_db() db.bindings[(agent_id, "extraction")]["model_deployment_id"] = None resolver = RegistryResolver(db, cache_ttl_seconds=60.0) with pytest.raises(InferenceError) as exc_info: await resolver.resolve_target(agent_id, "extraction") assert exc_info.value.category == InferenceErrorCategory.CAPABILITY_UNAVAILABLE assert "no deployment" in str(exc_info.value) # --------------------------------------------------------------------------- # RegistryResolver tests — auth_secret_ref preserved as-is # --------------------------------------------------------------------------- class TestResolverAuthPreservation: """Auth secret refs are preserved without resolution during caching.""" @pytest.mark.asyncio async def test_auth_secret_ref_preserved(self) -> None: """auth_secret_ref is kept as the reference string, not resolved.""" db, agent_id, endpoint_id, *_ = _build_test_db() db.endpoints[endpoint_id]["auth_secret_ref"] = "VLLM_API_KEY" resolver = RegistryResolver(db, cache_ttl_seconds=60.0) target = await resolver.resolve_target(agent_id, "extraction") # The reference is preserved as-is — no env var lookup assert target.auth_secret_ref == "VLLM_API_KEY" @pytest.mark.asyncio async def test_none_auth_secret_ref_preserved(self) -> None: """None auth_secret_ref is preserved as None.""" db, agent_id, endpoint_id, *_ = _build_test_db() db.endpoints[endpoint_id]["auth_secret_ref"] = None resolver = RegistryResolver(db, cache_ttl_seconds=60.0) target = await resolver.resolve_target(agent_id, "extraction") assert target.auth_secret_ref is None @pytest.mark.asyncio async def test_auth_not_resolved_from_env(self) -> None: """Even if env var exists, auth_secret_ref stays as reference string.""" db, agent_id, endpoint_id, *_ = _build_test_db() db.endpoints[endpoint_id]["auth_secret_ref"] = "MY_SECRET_KEY" resolver = RegistryResolver(db, cache_ttl_seconds=60.0) with patch.dict("os.environ", {"MY_SECRET_KEY": "actual-secret-value"}): target = await resolver.resolve_target(agent_id, "extraction") # Should be the reference, NOT the resolved value assert target.auth_secret_ref == "MY_SECRET_KEY" assert "actual-secret-value" not in str(target) # --------------------------------------------------------------------------- # RegistryResolver tests — deterministic resolution # --------------------------------------------------------------------------- class TestResolverDeterminism: """Given same DB state and same inputs, always returns same target.""" @pytest.mark.asyncio async def test_same_input_same_output(self) -> None: """Multiple resolutions with same state produce identical targets.""" db, agent_id, *_ = _build_test_db() resolver = RegistryResolver(db, cache_ttl_seconds=0.001) results = [] for _ in range(5): # Force re-resolution each time by expiring cache time.sleep(0.002) target = await resolver.resolve_target(agent_id, "extraction") results.append(target) # All results should be identical first = results[0] for r in results[1:]: assert r.endpoint_id == first.endpoint_id assert r.deployment_id == first.deployment_id assert r.protocol == first.protocol assert r.base_url == first.base_url assert r.model == first.model assert r.capabilities == first.capabilities assert r.auth_secret_ref == first.auth_secret_ref assert r.auth_scheme == first.auth_scheme assert r.extra_headers == first.extra_headers assert r.extra_body == first.extra_body assert r.context_window == first.context_window assert r.max_output_tokens == first.max_output_tokens @pytest.mark.asyncio async def test_separate_resolvers_same_result(self) -> None: """Two resolvers with same DB state produce identical targets.""" db, agent_id, *_ = _build_test_db() resolver1 = RegistryResolver(db, cache_ttl_seconds=60.0) resolver2 = RegistryResolver(db, cache_ttl_seconds=60.0) target1 = await resolver1.resolve_target(agent_id, "extraction") target2 = await resolver2.resolve_target(agent_id, "extraction") assert target1.endpoint_id == target2.endpoint_id assert target1.deployment_id == target2.deployment_id assert target1.protocol == target2.protocol assert target1.base_url == target2.base_url assert target1.model == target2.model assert target1.capabilities == target2.capabilities # --------------------------------------------------------------------------- # RegistryResolver tests — invalidation on revision/probe failure # --------------------------------------------------------------------------- class TestResolverInvalidationOnRevision: """Cache invalidation on revisions and failed probes.""" @pytest.mark.asyncio async def test_invalidation_after_endpoint_revision_change(self) -> None: """After endpoint revision changes, invalidation forces fresh data.""" db, agent_id, endpoint_id, *_ = _build_test_db() resolver = RegistryResolver(db, cache_ttl_seconds=60.0) # Initial resolution target1 = await resolver.resolve_target(agent_id, "extraction") assert target1.base_url == "http://vllm.stonks-oracle.svc:8000" # Simulate revision change (URL update) db.endpoints[endpoint_id]["base_url"] = "http://vllm-v2.stonks-oracle.svc:8000" db.endpoints[endpoint_id]["revision"] = 2 # Without invalidation, cache still returns old value target_cached = await resolver.resolve_target(agent_id, "extraction") assert target_cached.base_url == "http://vllm.stonks-oracle.svc:8000" # After invalidation, fresh data is fetched resolver.invalidate(endpoint_id) target2 = await resolver.resolve_target(agent_id, "extraction") assert target2.base_url == "http://vllm-v2.stonks-oracle.svc:8000" @pytest.mark.asyncio async def test_invalidation_simulates_failed_probe(self) -> None: """Simulated probe failure triggers invalidation and re-resolution.""" db, agent_id, endpoint_id, *_ = _build_test_db() resolver = RegistryResolver(db, cache_ttl_seconds=60.0) # Populate cache await resolver.resolve_target(agent_id, "extraction") initial_calls = db.call_count["get_active_binding"] # Simulate probe failure -> invalidate resolver.invalidate(endpoint_id) # Next resolution re-queries DB await resolver.resolve_target(agent_id, "extraction") assert db.call_count["get_active_binding"] == initial_calls + 1