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
3 changes: 3 additions & 0 deletions app/ai-service/config.py
Original file line number Diff line number Diff line change
Expand Up @@ -87,6 +87,9 @@ class Settings(BaseSettings):
# always appended so operators cannot accidentally expose themselves.
request_body_bypass_paths: str = ""

# Legacy route deprecation/retirement date (Sunset header)
legacy_retirement_date: str = "Wed, 01 Oct 2026 00:00:00 GMT"

# Verification artifact access settings
verification_artifacts_dir: str = "./artifacts/verification"
verification_artifact_url_ttl_seconds: int = 300
Expand Down
74 changes: 59 additions & 15 deletions app/ai-service/main.py
Original file line number Diff line number Diff line change
Expand Up @@ -18,6 +18,8 @@
from schemas.errors import ErrorDetail, ErrorEnvelope
import time
import metrics
import email.utils
from datetime import datetime, timezone

from slowapi import Limiter, _rate_limit_exceeded_handler
from slowapi.errors import RateLimitExceeded
Expand Down Expand Up @@ -352,6 +354,35 @@ class ProofOfLifeResponse(BaseModel):
app.add_exception_handler(RateLimitExceeded, _rate_limit_exceeded_handler)


def get_sunset_header_value() -> str:
val = settings.legacy_retirement_date
if not val:
return ""
val = val.strip()
# Try parsing various date formats to normalize to RFC 1123
for fmt in (
"%Y-%m-%d",
"%Y-%m-%dT%H:%M:%S",
"%Y-%m-%dT%H:%M:%SZ",
"%Y-%m-%d %H:%M:%S",
"%a, %d %b %Y %H:%M:%S %Z",
"%a, %d %b %Y %H:%M:%S",
):
try:
dt = datetime.strptime(val, fmt)
if dt.tzinfo is None:
dt = dt.replace(tzinfo=timezone.utc)
return email.utils.format_datetime(dt, usegmt=True)
except ValueError:
continue
try:
dt = email.utils.parsedate_to_datetime(val)
return email.utils.format_datetime(dt, usegmt=True)
except Exception:
pass
return val


@app.middleware("http")
async def legacy_redirect_middleware(request: Request, call_next):
"""
Expand All @@ -367,25 +398,38 @@ async def legacy_redirect_middleware(request: Request, call_next):
The /ai/metrics path is also excluded - it has no v1 equivalent.
"""
path = request.url.path
is_legacy = path.startswith("/ai/") and path != "/ai/metrics"

# Exact-match redirects
if path in _LEGACY_TO_V1:
target = _LEGACY_TO_V1[path]
if request.url.query:
target = f"{target}?{request.url.query}"
logger.debug(f"Legacy redirect: {path} -> {target}")
return RedirectResponse(url=target, status_code=308)

# Prefix-based redirects (parameterised routes)
for legacy_prefix, v1_prefix in _LEGACY_PREFIX_MAP:
if path.startswith(legacy_prefix):
target = v1_prefix + path[len(legacy_prefix) :]
response = None
if is_legacy:
# Exact-match redirects
if path in _LEGACY_TO_V1:
target = _LEGACY_TO_V1[path]
if request.url.query:
target = f"{target}?{request.url.query}"
logger.debug(f"Legacy prefix redirect: {path} -> {target}")
return RedirectResponse(url=target, status_code=308)
logger.debug(f"Legacy redirect: {path} -> {target}")
response = RedirectResponse(url=target, status_code=308)
else:
# Prefix-based redirects (parameterised routes)
for legacy_prefix, v1_prefix in _LEGACY_PREFIX_MAP:
if path.startswith(legacy_prefix):
target = v1_prefix + path[len(legacy_prefix) :]
if request.url.query:
target = f"{target}?{request.url.query}"
logger.debug(f"Legacy prefix redirect: {path} -> {target}")
response = RedirectResponse(url=target, status_code=308)
break

if response is None:
response = await call_next(request)

return await call_next(request)
if is_legacy:
sunset_val = get_sunset_header_value()
if sunset_val:
response.headers["Sunset"] = sunset_val
response.headers["Deprecation"] = "true"

return response


@app.middleware("http")
Expand Down
42 changes: 42 additions & 0 deletions app/ai-service/test_main.py
Original file line number Diff line number Diff line change
Expand Up @@ -226,3 +226,45 @@ def fake_verify_claim(aid_claim, supporting_evidence=None, context_factors=None,
data = response.json()
assert data["success"] is False
assert "all providers unavailable" in data["error"]


def test_legacy_routes_deprecation_headers(client, monkeypatch):
"""Test that all legacy /ai/* routes emit Sunset and Deprecation headers,
and preserve backwards compatibility (308 redirect or correct response)."""
# Temporarily set a specific retirement date for validation
monkeypatch.setattr(main.settings, "legacy_retirement_date", "2026-10-01")

legacy_redirect_routes = [
"/ai/inference",
"/ai/proof-of-life",
"/ai/anonymize",
"/ai/humanitarian/verify",
"/ai/status/test-task-id",
"/ai/task/test-task-id/cancel",
]

for route in legacy_redirect_routes:
# We test post for these legacy routes; middleware intercepts and returns 308
response = client.post(route, follow_redirects=False)
assert response.status_code == 308, f"Route {route} did not return 308 redirect"
assert response.headers.get("Deprecation") == "true", f"Route {route} missing Deprecation header"
assert response.headers.get("Sunset") == "Thu, 01 Oct 2026 00:00:00 GMT", f"Route {route} missing or invalid Sunset header"

# Also test /ai/ocr which is served directly and should not redirect, but still have headers
response = client.post("/ai/ocr", follow_redirects=False)
# /ai/ocr without files will return 422 or 400, but should still have the legacy headers
assert response.status_code in (400, 422), f"/ai/ocr returned unexpected status {response.status_code}"
assert response.headers.get("Deprecation") == "true", "/ai/ocr missing Deprecation header"
assert response.headers.get("Sunset") == "Thu, 01 Oct 2026 00:00:00 GMT", "/ai/ocr missing or invalid Sunset header"

# Test that non-legacy routes (like health, metrics, v1 paths) do NOT have the headers
non_legacy_routes = [
"/health",
"/ai/metrics",
"/v1/ai/inference",
]
for route in non_legacy_routes:
response = client.post(route, follow_redirects=False) if "inference" in route else client.get(route, follow_redirects=False)
assert "Deprecation" not in response.headers, f"Non-legacy route {route} has Deprecation header"
assert "Sunset" not in response.headers, f"Non-legacy route {route} has Sunset header"

119 changes: 76 additions & 43 deletions app/backend/src/common/security/security.module.ts
Original file line number Diff line number Diff line change
@@ -1,8 +1,10 @@
import { Module } from '@nestjs/common';
import { Logger, Module } from '@nestjs/common';
import type { CorsOptions } from '@nestjs/common/interfaces/external/cors-options.interface';
import { ConfigService } from '@nestjs/config';
import type { NextFunction, Request, RequestHandler, Response } from 'express';
import helmet, { HelmetOptions } from 'helmet';
import { RedisService } from '@liaoliaots/nestjs-redis';


const DEFAULT_ALLOWED_ORIGINS = [
'http://localhost:3000',
Expand Down Expand Up @@ -182,33 +184,24 @@ export const createCorsOriginValidator = (
};
};

export const createRateLimiter = (config: ConfigService): RequestHandler => {
export const createRateLimiter = (
config: ConfigService,
redisService?: RedisService,
): RequestHandler => {
const logger = new Logger('RateLimiter');

const windowMs = parseNumber(
config.get<string>('THROTTLE_TTL'),
config.get<string>('RATE_LIMIT_WINDOW_MS') ?? config.get<string>('THROTTLE_TTL'),
DEFAULT_RATE_LIMIT_WINDOW_MS,
);
const limit = parseNumber(
config.get<string>('API_RATE_LIMIT'),
config.get<string>('RATE_LIMIT_LIMIT') ?? config.get<string>('API_RATE_LIMIT'),
DEFAULT_RATE_LIMIT,
);

const store = new Map<string, { count: number; resetTimeMs: number }>();
let lastCleanupMs = 0;

const cleanupExpiredEntries = (now: number) => {
if (now - lastCleanupMs < windowMs) {
return;
}

lastCleanupMs = now;
for (const [key, entry] of store) {
if (entry.resetTimeMs <= now) {
store.delete(key);
}
}
};
const windowSeconds = Math.max(Math.ceil(windowMs / 1000), 1);

return (req: Request, res: Response, next: NextFunction) => {
return async (req: Request, res: Response, next: NextFunction) => {
if (isRateLimitExempt(req)) {
next();
return;
Expand All @@ -234,37 +227,77 @@ export const createRateLimiter = (config: ConfigService): RequestHandler => {
return;
}

const now = Date.now();
cleanupExpiredEntries(now);

const forwardedIp =
Array.isArray(req.ips) && req.ips.length > 0 ? req.ips[0] : undefined;
const key: string =
const key = `ratelimit:global:${
(typeof forwardedIp === 'string' ? forwardedIp : undefined) ??
(typeof req.ip === 'string' ? req.ip : undefined) ??
'unknown';
let entry = store.get(key);
if (!entry || entry.resetTimeMs <= now) {
entry = { count: 0, resetTimeMs: now + windowMs };
store.set(key, entry);
}
'unknown'
}`;

entry.count += 1;
const now = Date.now();
const minTimestamp = now - windowMs;
const uniqueMember = `${now}:${Math.random().toString(36).substring(2, 15)}`;

const remaining = Math.max(limit - entry.count, 0);
const resetSeconds = Math.max(
Math.ceil((entry.resetTimeMs - now) / 1000),
0,
);
try {
if (!redisService) {
throw new Error('RedisService is not configured/available');
}

res.setHeader('RateLimit-Limit', limit.toString());
res.setHeader('RateLimit-Remaining', remaining.toString());
res.setHeader('RateLimit-Reset', resetSeconds.toString());
const client = redisService.getOrThrow();

if (entry.count > limit) {
res.setHeader('Retry-After', resetSeconds.toString());
res.status(429).send('Too many requests, please try again later.');
return;
// Execute MULTI pipeline to keep header reads consistent
const multi = client.multi();
multi.zremrangebyscore(key, '-inf', minTimestamp);
multi.zadd(key, now, uniqueMember);
multi.zrange(key, 0, 0, 'WITHSCORES');
multi.zcard(key);
multi.expire(key, windowSeconds);

const results = await multi.exec();
if (!results) {
throw new Error('Redis multi transaction execution returned null');
}

const zrangeResult = results[2];
const zcardResult = results[3];

const zrangeRes = Array.isArray(zrangeResult) ? (zrangeResult[1] as string[]) : undefined;
const zcardRes = Array.isArray(zcardResult) ? (zcardResult[1] as number) : undefined;

const count = typeof zcardRes === 'number' ? zcardRes : 1;

// ZRANGE WITHSCORES returns: [member1, score1, member2, score2, ...]
// The oldest timestamp is the score of the first entry, i.e., index 1
let oldestTimestamp = now;
if (zrangeRes && zrangeRes.length >= 2) {
const parsed = Number(zrangeRes[1]);
if (!isNaN(parsed)) {
oldestTimestamp = parsed;
}
}

const remaining = Math.max(limit - count, 0);
const resetSeconds = Math.max(
Math.ceil((oldestTimestamp + windowMs - now) / 1000),
0,
);

res.setHeader('RateLimit-Limit', limit.toString());
res.setHeader('RateLimit-Remaining', remaining.toString());
res.setHeader('RateLimit-Reset', resetSeconds.toString());

if (count > limit) {
res.setHeader('Retry-After', resetSeconds.toString());
res.status(429).send('Too many requests, please try again later.');
return;
}
} catch (err) {
logger.warn(
`Redis rate limiter failed, failing open: ${
err instanceof Error ? err.message : String(err)
}`,
);
}

next();
Expand Down
3 changes: 2 additions & 1 deletion app/backend/src/main.ts
Original file line number Diff line number Diff line change
Expand Up @@ -18,6 +18,7 @@ import {
createHelmetMiddleware,
createRateLimiter,
} from './common/security/security.module';
import { RedisService } from '@liaoliaots/nestjs-redis';

async function bootstrap() {
// Load environment variables
Expand Down Expand Up @@ -49,7 +50,7 @@ async function bootstrap() {
app.use(createHelmetMiddleware(configService));
app.use(createCorsOriginValidator(configService));
app.enableCors(buildCorsOptions(configService));
app.use(createRateLimiter(configService));
app.use(createRateLimiter(configService, app.get(RedisService)));
app.use(compression());

// Global prefix
Expand Down
Loading
Loading