Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
26 changes: 25 additions & 1 deletion backend/app/services/bridge_sync/cache_refresh.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down Expand Up @@ -388,14 +411,15 @@ 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",
timeout=request_timeout,
) 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:
Expand Down
56 changes: 56 additions & 0 deletions backend/tests/services/test_bridge_cache_refresh.py
Original file line number Diff line number Diff line change
Expand Up @@ -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()
Expand Down
Loading