diff --git a/backend/app/services/bridge_sync/cache_refresh.py b/backend/app/services/bridge_sync/cache_refresh.py index 92d92ee0..5dc6ac36 100644 --- a/backend/app/services/bridge_sync/cache_refresh.py +++ b/backend/app/services/bridge_sync/cache_refresh.py @@ -39,6 +39,29 @@ def _env_int(name: str, default: int) -> int: return default +def _configured_server_api_key() -> str | None: + """Return a server API key accepted by the normal rate-limit middleware.""" + for env_name in ("BTAA_GEOSPATIAL_API_KEY", "BTAA_GEOSPATIAL_API_KEYS"): + for candidate in os.getenv(env_name, "").split(","): + key = candidate.strip() + if key: + return key + return None + + +def _rewarm_request_headers() -> dict[str, str]: + headers = {"Accept": "application/json"} + api_key = _configured_server_api_key() + if api_key: + headers["X-API-Key"] = api_key + else: + logger.warning( + "Bridge cache rewarm has no configured server API key; " + "requests may be subject to the anonymous rate limit" + ) + return headers + + def _positive_env_int(name: str, default: int) -> int: value = _env_int(name, default) if value < 1: @@ -388,6 +411,7 @@ def add_warm_path(path: str | None) -> None: from app.main import app transport = httpx.ASGITransport(app=app) + request_headers = _rewarm_request_headers() async with httpx.AsyncClient( transport=transport, base_url="http://bridge-cache-refresh.local", @@ -395,7 +419,7 @@ def add_warm_path(path: str | None) -> None: ) as client: for path in warm_paths: try: - response = await client.get(path, headers={"Accept": "application/json"}) + response = await client.get(path, headers=request_headers) if 200 <= response.status_code < 300: warmed += 1 else: diff --git a/backend/tests/services/test_bridge_cache_refresh.py b/backend/tests/services/test_bridge_cache_refresh.py index 25abd0a6..4e61987f 100644 --- a/backend/tests/services/test_bridge_cache_refresh.py +++ b/backend/tests/services/test_bridge_cache_refresh.py @@ -122,6 +122,62 @@ async def fake_warm_assets(resource_ids): ] +@pytest.mark.asyncio +async def test_cache_rewarm_uses_server_key_beyond_anonymous_rate_limit(monkeypatch): + class FakeResponse: + def __init__(self, status_code): + self.status_code = status_code + + class RateLimitedAsyncClient: + def __init__(self): + self.calls = [] + + async def __aenter__(self): + return self + + async def __aexit__(self, exc_type, exc, traceback): + return False + + async def get(self, path, *, headers): + self.calls.append((path, dict(headers))) + authenticated = headers.get("X-API-Key") == "server-api-key" + status_code = 200 if authenticated or len(self.calls) <= 10 else 429 + return FakeResponse(status_code) + + fake_cache = FakeCacheService() + fake_cache.cached_records_for_tags = AsyncMock(return_value=[]) + fake_client = RateLimitedAsyncClient() + resource_ids = [f"resource-{index}" for index in range(12)] + + monkeypatch.setattr(cache_refresh, "ENDPOINT_CACHE", True) + monkeypatch.setenv("BRIDGE_CACHE_REFRESH_ENABLED", "true") + monkeypatch.setenv("BTAA_GEOSPATIAL_API_KEY", "server-api-key") + + with ( + patch.object(cache_refresh, "CacheService", return_value=fake_cache), + patch.object( + cache_refresh, + "delete_resource_representations", + new=AsyncMock(return_value={"durable_deleted": True, "redis_deleted": 0}), + ), + patch.object(cache_refresh.httpx, "ASGITransport", return_value=object()), + patch.object(cache_refresh.httpx, "AsyncClient", return_value=fake_client), + ): + stats = await cache_refresh.refresh_cache_for_changed_resources( + resource_ids, + warm_generated_assets=False, + ) + + assert stats["warm_urls"] == 12 + assert stats["warmed"] == 12 + assert stats["errors"] == 0 + assert len(fake_client.calls) == 12 + assert all( + headers == {"Accept": "application/json", "X-API-Key": "server-api-key"} + for _path, headers in fake_client.calls + ) + + @pytest.mark.asyncio async def test_lightweight_refresh_invalidates_without_rewarming_or_generating(monkeypatch): fake_cache = FakeCacheService()