diff --git a/app/ai-service/config.py b/app/ai-service/config.py index 67367e49..6ba1bc64 100644 --- a/app/ai-service/config.py +++ b/app/ai-service/config.py @@ -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 diff --git a/app/ai-service/main.py b/app/ai-service/main.py index b067e438..f26632ed 100644 --- a/app/ai-service/main.py +++ b/app/ai-service/main.py @@ -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 @@ -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): """ @@ -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") diff --git a/app/ai-service/test_main.py b/app/ai-service/test_main.py index e0f83c56..1180c792 100644 --- a/app/ai-service/test_main.py +++ b/app/ai-service/test_main.py @@ -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" + diff --git a/app/backend/src/common/security/security.module.ts b/app/backend/src/common/security/security.module.ts index f0961b12..1e2ddc46 100644 --- a/app/backend/src/common/security/security.module.ts +++ b/app/backend/src/common/security/security.module.ts @@ -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', @@ -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('THROTTLE_TTL'), + config.get('RATE_LIMIT_WINDOW_MS') ?? config.get('THROTTLE_TTL'), DEFAULT_RATE_LIMIT_WINDOW_MS, ); const limit = parseNumber( - config.get('API_RATE_LIMIT'), + config.get('RATE_LIMIT_LIMIT') ?? config.get('API_RATE_LIMIT'), DEFAULT_RATE_LIMIT, ); - const store = new Map(); - 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; @@ -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(); diff --git a/app/backend/src/main.ts b/app/backend/src/main.ts index 9fa25572..3e5ea7c5 100644 --- a/app/backend/src/main.ts +++ b/app/backend/src/main.ts @@ -18,6 +18,7 @@ import { createHelmetMiddleware, createRateLimiter, } from './common/security/security.module'; +import { RedisService } from '@liaoliaots/nestjs-redis'; async function bootstrap() { // Load environment variables @@ -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 diff --git a/app/backend/test/security.e2e-spec.ts b/app/backend/test/security.e2e-spec.ts index c76c8e12..129bb504 100644 --- a/app/backend/test/security.e2e-spec.ts +++ b/app/backend/test/security.e2e-spec.ts @@ -1,4 +1,4 @@ -import { INestApplication, VersioningType } from '@nestjs/common'; +import { Logger, INestApplication, VersioningType } from '@nestjs/common'; import { ConfigService } from '@nestjs/config'; import { Test, TestingModule } from '@nestjs/testing'; import { DocumentBuilder, SwaggerModule } from '@nestjs/swagger'; @@ -10,6 +10,8 @@ import { createHelmetMiddleware, createRateLimiter, } from '../src/common/security/security.module'; +import { RedisService } from '@liaoliaots/nestjs-redis'; +import RedisMock from 'ioredis-mock'; type TestAppOptions = { enableDocs: boolean; @@ -41,7 +43,7 @@ const createTestApp = async ({ enableDocs }: TestAppOptions) => { app.use(createHelmetMiddleware(configService)); app.use(createCorsOriginValidator(configService)); app.enableCors(buildCorsOptions(configService)); - app.use(createRateLimiter(configService)); + app.use(createRateLimiter(configService, app.get(RedisService))); if (enableDocs) { const swaggerConfig = new DocumentBuilder() @@ -173,6 +175,7 @@ describe('Security (e2e)', () => { let now = initialNow; let nowSpy: jest.SpyInstance; let rateLimitApp: INestApplication; + let mockRedisClient: any; beforeEach(async () => { process.env.API_RATE_LIMIT = '2'; @@ -183,6 +186,10 @@ describe('Security (e2e)', () => { now = initialNow; nowSpy = jest.spyOn(Date, 'now').mockImplementation(() => now); rateLimitApp = await createTestApp({ enableDocs: true }); + + mockRedisClient = new RedisMock(); + const redisService = rateLimitApp.get(RedisService); + jest.spyOn(redisService, 'getOrThrow').mockReturnValue(mockRedisClient); }); afterEach(async () => { @@ -193,6 +200,8 @@ describe('Security (e2e)', () => { process.env.THROTTLE_TTL = '60000'; process.env.CORS_ORIGINS = 'http://localhost:3000'; process.env.CORS_ALLOW_CREDENTIALS = 'false'; + delete process.env.RATE_LIMIT_LIMIT; + delete process.env.RATE_LIMIT_WINDOW_MS; }); it('should rate limit, include retry headers, and reset after the window passes', async () => { @@ -231,6 +240,59 @@ describe('Security (e2e)', () => { expect(response.status).toBe(200); } }); + + it('should rate limit 100 hits in 1 s => 80+ return 429, and include correct headers', async () => { + // Create a specific application instance configured for 20 req/s + process.env.RATE_LIMIT_LIMIT = '20'; + process.env.RATE_LIMIT_WINDOW_MS = '1000'; + + const appInstance = await createTestApp({ enableDocs: false }); + const redisService = appInstance.get(RedisService); + const testMockRedis = new RedisMock(); + jest.spyOn(redisService, 'getOrThrow').mockReturnValue(testMockRedis as any); + + const server = appInstance.getHttpServer(); + const results: any[] = []; + + for (let i = 0; i < 100; i += 1) { + results.push(request(server).get('/api/v1/')); + } + + const responses = await Promise.all(results); + const count429 = responses.filter(r => r.status === 429).length; + + expect(count429).toBeGreaterThanOrEqual(80); + + const rateLimitedResponse = responses.find(r => r.status === 429); + expect(rateLimitedResponse).toBeDefined(); + expect(rateLimitedResponse.headers['ratelimit-limit']).toBe('20'); + expect(rateLimitedResponse.headers['ratelimit-remaining']).toBeDefined(); + expect(rateLimitedResponse.headers['ratelimit-reset']).toBeDefined(); + expect(rateLimitedResponse.headers['retry-after']).toBeDefined(); + + await appInstance.close(); + }); + + it('should fail open with a WARN log, not 500, when Redis is down', async () => { + const appInstance = await createTestApp({ enableDocs: false }); + const redisService = appInstance.get(RedisService); + + jest.spyOn(redisService, 'getOrThrow').mockImplementation(() => { + throw new Error('Redis connection down'); + }); + + const warnSpy = jest.spyOn(Logger.prototype, 'warn').mockImplementation(() => {}); + + const server = appInstance.getHttpServer(); + const response = await request(server).get('/api/v1/'); + + expect(response.status).not.toBe(500); + expect(response.status).not.toBe(429); + expect(warnSpy).toHaveBeenCalled(); + + warnSpy.mockRestore(); + await appInstance.close(); + }); }); describe('Docs Endpoint', () => { diff --git a/pnpm-lock.yaml b/pnpm-lock.yaml index fdd7d856..e393c248 100644 --- a/pnpm-lock.yaml +++ b/pnpm-lock.yaml @@ -136,6 +136,9 @@ importers: class-validator: specifier: ^0.14.3 version: 0.14.4 + compression: + specifier: 1.7.5 + version: 1.7.5 dotenv: specifier: ^17.2.3 version: 17.4.2 @@ -185,6 +188,9 @@ importers: '@nestjs/testing': specifier: ^11.0.1 version: 11.1.17(@nestjs/common@11.1.27(class-transformer@0.5.1)(class-validator@0.14.4)(reflect-metadata@0.2.2)(rxjs@7.8.2))(@nestjs/core@11.1.17(@nestjs/common@11.1.27(class-transformer@0.5.1)(class-validator@0.14.4)(reflect-metadata@0.2.2)(rxjs@7.8.2))(@nestjs/platform-express@11.1.17)(reflect-metadata@0.2.2)(rxjs@7.8.2))(@nestjs/platform-express@11.1.17(@nestjs/common@11.1.27(class-transformer@0.5.1)(class-validator@0.14.4)(reflect-metadata@0.2.2)(rxjs@7.8.2))(@nestjs/core@11.1.17)) + '@types/compression': + specifier: 1.7.5 + version: 1.7.5 '@types/express': specifier: ^5.0.0 version: 5.0.6 @@ -3242,6 +3248,9 @@ packages: '@types/body-parser@1.19.6': resolution: {integrity: sha512-HLFeCYgz89uk22N5Qg3dvGvsv46B8GLvKKo1zKG4NybA8U2DiEO3w9lqGg29t/tfLRJpJ6iQxnVw4OnB7MoM9g==} + '@types/compression@1.7.5': + resolution: {integrity: sha512-AAQvK5pxMpaT+nDvhHrsBhLSYG5yQdtkaJE1WYieSNY2mVFKAgmU4ks65rkZD5oqnGCFLyQpUr1CqI4DmUMyDg==} + '@types/connect@3.4.38': resolution: {integrity: sha512-K6uROf1LD88uDQqJCktA4yzL1YYAK6NgfsI0v/mTgyPKWsX1CnJ0XPSDhViejru1GcRkLWb8RlzFYJRqGUbaug==} @@ -4370,8 +4379,8 @@ packages: resolution: {integrity: sha512-AF3r7P5dWxL8MxyITRMlORQNaOA2IkAFaTr4k7BUumjPtRpGDTZpl0Pb1XCO6JeDCBdp126Cgs9sMxqSjgYyRg==} engines: {node: '>= 0.6'} - compression@1.8.1: - resolution: {integrity: sha512-9mAqGPHLakhCLeNyxPkK4xVo746zQ/czLH1Ky+vkitMnWfWZps8r0qXuwhwizagCRttsL4lfG4pIOvaWLpAP0w==} + compression@1.7.5: + resolution: {integrity: sha512-bQJ0YRck5ak3LgtnpKkiabX5pNF7tMUh1BSy2ZBOTh0Dim0BUu6aPPwByIns6/A5Prh8PufSPerMDUklpzes2Q==} engines: {node: '>= 0.8.0'} concat-map@0.0.1: @@ -6973,8 +6982,8 @@ packages: resolution: {integrity: sha512-oVlzkg3ENAhCk2zdv7IJwd/QUD4z2RxRwpkcGY8psCVcCYZNq4wYnVWALHM+brtuJjePWiYF/ClmuDr8Ch5+kg==} engines: {node: '>= 0.8'} - on-headers@1.1.0: - resolution: {integrity: sha512-737ZY3yNnXy37FHkQxPzt4UZ2UWPWiCZWLvFZ4fu5cueciegX0zGPnrlY6bwRg4FdQOe9YU8MkmJwGhoMybl8A==} + on-headers@1.0.2: + resolution: {integrity: sha512-pZAE+FJLoyITytdqK0U5s+FIpjN0JP3OzFi/u8Rx+EV5/W+JTWGXG8xFzevE7AjBfDqHv/8vL8qQsIhHnqRkrA==} engines: {node: '>= 0.8'} once@1.4.0: @@ -9792,7 +9801,7 @@ snapshots: bplist-parser: 0.3.2 chalk: 4.1.2 ci-info: 3.9.0 - compression: 1.8.1 + compression: 1.7.5 connect: 3.7.0 debug: 4.4.3(supports-color@10.2.2) env-editor: 0.4.2 @@ -12192,6 +12201,10 @@ snapshots: '@types/connect': 3.4.38 '@types/node': 25.9.1 + '@types/compression@1.7.5': + dependencies: + '@types/express': 5.0.6 + '@types/connect@3.4.38': dependencies: '@types/node': 25.9.1 @@ -13661,13 +13674,13 @@ snapshots: dependencies: mime-db: 1.54.0 - compression@1.8.1: + compression@1.7.5: dependencies: bytes: 3.1.2 compressible: 2.0.18 debug: 2.6.9 negotiator: 0.6.4 - on-headers: 1.1.0 + on-headers: 1.0.2 safe-buffer: 5.2.1 vary: 1.1.2 transitivePeerDependencies: @@ -14157,7 +14170,7 @@ snapshots: tinyglobby: 0.2.15 unrs-resolver: 1.11.1 optionalDependencies: - eslint-plugin-import: 2.32.0(@typescript-eslint/parser@8.57.1(eslint@9.39.4(jiti@2.6.1))(typescript@5.9.3))(eslint-import-resolver-typescript@3.10.1)(eslint@9.39.4(jiti@2.6.1)) + eslint-plugin-import: 2.32.0(eslint-import-resolver-typescript@3.10.1)(eslint@9.39.4(jiti@2.6.1)) transitivePeerDependencies: - supports-color @@ -17178,7 +17191,7 @@ snapshots: dependencies: ee-first: 1.1.1 - on-headers@1.1.0: {} + on-headers@1.0.2: {} once@1.4.0: dependencies: