Implement response caching with configurable settings in API gateway

This commit is contained in:
2026-05-31 23:43:06 +08:00
parent a9ae468817
commit 3270568654
7 changed files with 201 additions and 5 deletions
+10
View File
@@ -56,6 +56,9 @@ settings:
version: string
description: string
external_openapi_file: path/to/openapi.json
cache_enabled: false
cache_ttl_seconds: 30
cache_max_entries: 500
require_api_key: false
api_keys_file: config/api_keys.json
admin_api_key_env: ADMIN_API_KEY
@@ -100,6 +103,13 @@ Optional admin auth:
- Clients then must send `x-api-key: <key>` on proxied route requests.
- Create keys from the admin portal or `POST /admin/api-keys`.
## Response Cache
- Enable cache with `settings.cache_enabled: true`.
- Configure TTL with `settings.cache_ttl_seconds`.
- Configure memory bound with `settings.cache_max_entries`.
- Cache currently applies to proxied `GET`/`HEAD` responses with `2xx` status.
- Gateway adds `x-gateway-cache: MISS` for first fetch and `x-gateway-cache: HIT` for cached responses.
## Scaling and Configurability Guidance
- Add routes through config, not code, for repeatable deployments.
- Split large configurations into environment-specific files and set `GATEWAY_CONFIG_PATH`.
+3
View File
@@ -3,6 +3,9 @@ settings:
version: 0.1.0
description: Configurable gateway for OpenAPI-based upstream services
external_openapi_file: openapi_json/lta_datamall_openapi_v0-1-1.json
cache_enabled: true
cache_ttl_seconds: 30
cache_max_entries: 500
require_api_key: false
api_keys_file: config/api_keys.json
admin_api_key_env: ADMIN_API_KEY
+9
View File
@@ -15,6 +15,7 @@ from pydantic import BaseModel
from .admin import get_admin_assets_dir, render_admin_portal, require_admin_access
from .api_keys import ApiKeyStore
from .cache import ResponseCache
from .config import GatewayConfig, RouteConfig, load_gateway_config
from .proxy import proxy_request
@@ -165,6 +166,10 @@ def _register_management_routes(app: FastAPI) -> None:
request.app.state.gateway_config = reloaded
request.app.state.api_key_store = ApiKeyStore(Path(reloaded.settings.api_keys_file))
request.app.state.api_key_store.ensure_exists()
request.app.state.response_cache = ResponseCache(
ttl_seconds=reloaded.settings.cache_ttl_seconds,
max_entries=reloaded.settings.cache_max_entries,
)
request.app.state.admin_api_key = os.getenv(reloaded.settings.admin_api_key_env, "")
return {"status": "saved"}
@@ -202,6 +207,10 @@ def create_app(config_path: str | None = None) -> FastAPI:
app.state.http_client = httpx.AsyncClient()
app.state.api_key_store = ApiKeyStore(Path(gateway_config.settings.api_keys_file))
app.state.api_key_store.ensure_exists()
app.state.response_cache = ResponseCache(
ttl_seconds=gateway_config.settings.cache_ttl_seconds,
max_entries=gateway_config.settings.cache_max_entries,
)
app.state.admin_api_key = os.getenv(gateway_config.settings.admin_api_key_env, "")
try:
yield
+66
View File
@@ -0,0 +1,66 @@
from __future__ import annotations
from collections import OrderedDict
from dataclasses import dataclass
import time
@dataclass
class CachedResponse:
status_code: int
headers: dict[str, str]
body: bytes
media_type: str | None
expires_at: float
class ResponseCache:
def __init__(self, *, ttl_seconds: float, max_entries: int) -> None:
self.ttl_seconds = ttl_seconds
self.max_entries = max_entries
self._entries: OrderedDict[str, CachedResponse] = OrderedDict()
def _purge_expired(self) -> None:
now = time.monotonic()
expired = [key for key, value in self._entries.items() if value.expires_at <= now]
for key in expired:
self._entries.pop(key, None)
def get(self, key: str) -> CachedResponse | None:
self._purge_expired()
entry = self._entries.get(key)
if entry is None:
return None
self._entries.move_to_end(key)
return entry
def set(
self,
key: str,
*,
status_code: int,
headers: dict[str, str],
body: bytes,
media_type: str | None,
) -> None:
if self.ttl_seconds <= 0 or self.max_entries <= 0:
return
self._purge_expired()
if key in self._entries:
self._entries.pop(key, None)
while len(self._entries) >= self.max_entries:
self._entries.popitem(last=False)
self._entries[key] = CachedResponse(
status_code=status_code,
headers=dict(headers),
body=bytes(body),
media_type=media_type,
expires_at=time.monotonic() + self.ttl_seconds,
)
def clear(self) -> None:
self._entries.clear()
+3
View File
@@ -12,6 +12,9 @@ class GatewaySettings(BaseModel):
version: str = "0.1.0"
description: str = "Configurable API gateway"
external_openapi_file: str | None = None
cache_enabled: bool = False
cache_ttl_seconds: float = Field(default=30.0, ge=1.0, le=3600.0)
cache_max_entries: int = Field(default=500, ge=1, le=10000)
require_api_key: bool = False
api_keys_file: str = "config/api_keys.json"
admin_api_key_env: str = "ADMIN_API_KEY"
+54 -2
View File
@@ -21,6 +21,23 @@ HOP_BY_HOP_HEADERS = {
}
def _cache_key(request: Request, target_url: str) -> str:
params = "&".join(f"{k}={v}" for k, v in request.query_params.multi_items())
vary_accept = request.headers.get("accept", "")
vary_auth = request.headers.get("authorization", "")
return f"{request.method}|{target_url}|{params}|{vary_accept}|{vary_auth}"
def _get_cache_from_request(request: Request):
app = request.scope.get("app")
if app is None:
return None
state = getattr(app, "state", None)
if state is None:
return None
return getattr(state, "response_cache", None)
def build_target_path(request: Request, route: RouteConfig) -> str:
incoming_path = request.url.path
@@ -68,6 +85,26 @@ async def proxy_request(
target_path = build_target_path(request, route)
target_url = urljoin(str(upstream.base_url), target_path.lstrip("/"))
cache = None
cache_key = ""
cache_enabled = (
config.settings.cache_enabled and request.method in {"GET", "HEAD"}
)
if cache_enabled:
cache = _get_cache_from_request(request)
if cache is not None:
cache_key = _cache_key(request, target_url)
cached = cache.get(cache_key)
if cached is not None:
headers = dict(cached.headers)
headers["x-gateway-cache"] = "HIT"
return Response(
content=cached.body,
status_code=cached.status_code,
headers=headers,
media_type=cached.media_type,
)
try:
upstream_response = await client.request(
method=request.method,
@@ -88,9 +125,24 @@ async def proxy_request(
media_type="application/json",
)
return Response(
headers = _response_headers(upstream_response)
if cache_enabled:
headers["x-gateway-cache"] = "MISS"
response = Response(
content=upstream_response.content,
status_code=upstream_response.status_code,
headers=_response_headers(upstream_response),
headers=headers,
media_type=upstream_response.headers.get("content-type"),
)
if cache_enabled and cache is not None and 200 <= upstream_response.status_code < 300:
cache.set(
cache_key,
status_code=upstream_response.status_code,
headers=_response_headers(upstream_response),
body=upstream_response.content,
media_type=upstream_response.headers.get("content-type"),
)
return response
+56 -3
View File
@@ -1,21 +1,31 @@
from __future__ import annotations
from types import SimpleNamespace
import httpx
import pytest
from starlette.requests import Request
from gateway_framework.config import GatewayConfig, RouteConfig, UpstreamConfig
from gateway_framework.cache import ResponseCache
from gateway_framework.config import GatewayConfig, GatewaySettings, RouteConfig, UpstreamConfig
from gateway_framework.proxy import build_target_path, proxy_request
def _make_request(path: str, query: str = "", headers: list[tuple[str, str]] | None = None) -> Request:
def _make_request(
path: str,
query: str = "",
headers: list[tuple[str, str]] | None = None,
*,
app: object | None = None,
method: str = "GET",
) -> Request:
encoded_headers = [
(key.lower().encode("utf-8"), value.encode("utf-8")) for key, value in (headers or [])
]
scope = {
"type": "http",
"http_version": "1.1",
"method": "GET",
"method": method,
"scheme": "http",
"path": path,
"query_string": query.encode("utf-8"),
@@ -23,6 +33,8 @@ def _make_request(path: str, query: str = "", headers: list[tuple[str, str]] | N
"client": ("127.0.0.1", 51234),
"server": ("gateway.local", 8000),
}
if app is not None:
scope["app"] = app
async def receive() -> dict[str, object]:
return {"type": "http.request", "body": b"", "more_body": False}
@@ -77,6 +89,47 @@ async def test_proxy_request_forwards_to_upstream() -> None:
assert seen["x_forwarded_host"] == "gateway.local"
@pytest.mark.anyio
async def test_proxy_request_cache_hit_skips_second_upstream_call() -> None:
calls = {"count": 0}
def handler(req: httpx.Request) -> httpx.Response:
calls["count"] += 1
return httpx.Response(
status_code=200,
content=b'{"cached":true}',
headers={"content-type": "application/json"},
)
config = GatewayConfig(
settings=GatewaySettings(cache_enabled=True, cache_ttl_seconds=60, cache_max_entries=100),
upstreams={"demo": UpstreamConfig(base_url="https://example.com/", timeout_seconds=5)},
routes=[],
)
route = RouteConfig(path="/api/v1/bus", methods=["GET"], upstream="demo")
fake_app = SimpleNamespace(state=SimpleNamespace(response_cache=ResponseCache(ttl_seconds=60, max_entries=10)))
async with httpx.AsyncClient(transport=httpx.MockTransport(handler)) as client:
first = await proxy_request(
request=_make_request("/api/v1/bus", query="a=1", app=fake_app),
client=client,
config=config,
route=route,
)
second = await proxy_request(
request=_make_request("/api/v1/bus", query="a=1", app=fake_app),
client=client,
config=config,
route=route,
)
assert first.status_code == 200
assert first.headers.get("x-gateway-cache") == "MISS"
assert second.status_code == 200
assert second.headers.get("x-gateway-cache") == "HIT"
assert calls["count"] == 1
@pytest.mark.anyio
async def test_proxy_request_returns_502_when_upstream_unreachable() -> None:
def handler(req: httpx.Request) -> httpx.Response: