""" Error logging middleware for capturing and storing API errors. Request logs are buffered in an in-memory async queue and flushed in batch every second (or when the buffer reaches 500 entries) by a background task started at app startup via start_request_log_flusher(). This avoids one DB write per request at high throughput (1 M+ req/day) and eliminates DB lock contention as a contributor to p95 tail latency. Error logs (4xx/5xx) are still written synchronously because they are relatively rare and need accurate timing. Pure ASGI implementation — does NOT subclass BaseHTTPMiddleware. BaseHTTPMiddleware spawns an extra asyncio task per request via call_next(), doubling event-loop task counts under load and causing health-check timeouts. A raw ASGI __call__ wraps the downstream send callable instead, staying in the same task throughout. """ import asyncio import json import time import traceback import uuid from datetime import datetime, timedelta, timezone from typing import Optional from urllib.parse import parse_qs from starlette.responses import JSONResponse from starlette.types import ASGIApp, Receive, Scope, Send, Message from app.models.error_log import ErrorLog from app.models.request_log import RequestLog import logging logger = logging.getLogger(__name__) # --------------------------------------------------------------------------- # Async queue for request log entries # --------------------------------------------------------------------------- _REQUEST_LOG_QUEUE: asyncio.Queue = asyncio.Queue(maxsize=10_000) _FLUSH_INTERVAL_SECONDS: float = 2.0 # flush at most every 2 s _FLUSH_BATCH_SIZE: int = 500 # or when 500 entries are queued _LOG_RETENTION_DAYS: int = 7 # delete request_logs older than this _CLEANUP_INTERVAL_HOURS: int = 1 # run cleanup every N hours async def _flush_request_logs(entries: list) -> None: """Batch-insert a list of RequestLog dicts into the database.""" if not entries: return try: from app.core.database import AsyncSessionLocal async with AsyncSessionLocal() as db: for entry in entries: db.add(RequestLog(**entry)) await db.commit() except Exception as e: logger.error(f"Failed to flush request log batch ({len(entries)} entries): {e}") async def _request_log_flusher() -> None: """Background coroutine: drain the queue and batch-insert into DB.""" while True: entries: list = [] try: # Wait for at least one entry first = await asyncio.wait_for( _REQUEST_LOG_QUEUE.get(), timeout=_FLUSH_INTERVAL_SECONDS ) entries.append(first) _REQUEST_LOG_QUEUE.task_done() # Drain up to FLUSH_BATCH_SIZE without blocking while len(entries) < _FLUSH_BATCH_SIZE: try: item = _REQUEST_LOG_QUEUE.get_nowait() entries.append(item) _REQUEST_LOG_QUEUE.task_done() except asyncio.QueueEmpty: break except asyncio.TimeoutError: pass # Nothing queued in the last interval — loop again if entries: await _flush_request_logs(entries) async def _request_log_cleanup() -> None: """Background coroutine: delete request_logs older than _LOG_RETENTION_DAYS.""" while True: await asyncio.sleep(_CLEANUP_INTERVAL_HOURS * 3600) try: cutoff = datetime.now(timezone.utc) - timedelta(days=_LOG_RETENTION_DAYS) from app.core.database import AsyncSessionLocal from sqlalchemy import text async with AsyncSessionLocal() as db: result = await db.execute( text("DELETE FROM request_logs WHERE created_at < :cutoff"), {"cutoff": cutoff}, ) deleted = result.rowcount await db.commit() if deleted: logger.info(f"request_logs cleanup: deleted {deleted} rows older than {_LOG_RETENTION_DAYS}d") except Exception as e: logger.error(f"request_logs cleanup failed: {e}") def start_request_log_flusher() -> None: """Schedule the background flusher coroutine. Call once at app startup.""" asyncio.ensure_future(_request_log_flusher()) asyncio.ensure_future(_request_log_cleanup()) logger.info("Request log flusher background task started") class ErrorLoggingMiddleware: """Pure ASGI middleware: logs all requests and 4xx/5xx errors to the database. Unlike a BaseHTTPMiddleware subclass, this implementation wraps the ASGI `send` callable directly. Every request is handled in a single asyncio task — no call_next() extra task, no event-loop task doubling. """ def __init__(self, app: ASGIApp) -> None: self.app = app async def __call__(self, scope: Scope, receive: Receive, send: Send) -> None: # Pass non-HTTP scopes (lifespan, websocket) straight through. if scope["type"] != "http": await self.app(scope, receive, send) return request_id = str(uuid.uuid4())[:8] start_time = time.time() request_info = self._extract_request_info(scope) # Mutable state shared with the send wrapper (nonlocal in closure). status_code: list[int] = [200] response_headers: list = [] body_chunks: list[bytes] = [] is_error: list[bool] = [False] async def send_wrapper(message: Message) -> None: if message["type"] == "http.response.start": status_code[0] = message["status"] is_error[0] = status_code[0] >= 400 # Inject X-Request-ID into the response headers. headers = list(message.get("headers", [])) headers.append((b"x-request-id", request_id.encode())) response_headers.extend(headers) await send({**message, "headers": headers}) elif message["type"] == "http.response.body": if is_error[0]: body_chunks.append(message.get("body", b"")) await send(message) else: await send(message) try: await self.app(scope, receive, send_wrapper) except Exception as exc: response_time_ms = (time.time() - start_time) * 1000 await self._log_error( request_id=request_id, request_info=request_info, status_code=500, error_type="INTERNAL_SERVER_ERROR", error_message=str(exc), stack_trace=traceback.format_exc(), response_time_ms=response_time_ms, ) error_response = JSONResponse( status_code=500, content={ "error_type": "INTERNAL_SERVER_ERROR", "message": "An unexpected error occurred", "request_id": request_id, "timestamp": datetime.now(timezone.utc).isoformat(), }, headers={"X-Request-ID": request_id}, ) await error_response(scope, receive, send) return # --- Post-response logging (runs after all body chunks are sent) --- response_time_ms = (time.time() - start_time) * 1000 if is_error[0]: response_body = b"".join(body_chunks) error_detail = None try: if response_body: error_detail = json.loads(response_body.decode("utf-8")) except Exception: try: error_detail = {"raw_response": response_body.decode("utf-8")} except Exception: error_detail = {"raw_response": str(response_body)} await self._log_error( request_id=request_id, request_info=request_info, status_code=status_code[0], error_detail=error_detail, response_time_ms=response_time_ms, ) # Extract X-Data-Source from response headers if present. data_source = None for name, value in response_headers: if name == b"x-data-source": data_source = value.decode() break response_size = sum(len(c) for c in body_chunks) if body_chunks else None await self._log_request( request_id=request_id, request_info=request_info, status_code=status_code[0], response_time_ms=response_time_ms, response_size=response_size, data_source=data_source, ) # ------------------------------------------------------------------ # Helpers # ------------------------------------------------------------------ def _extract_request_info(self, scope: Scope) -> dict: """Extract minimal request info from the ASGI scope (sync, no body read).""" headers_raw = dict(scope.get("headers", [])) user_agent = headers_raw.get(b"user-agent", b"").decode("latin-1") or None path = scope.get("path", "") qs_bytes = scope.get("query_string", b"") qs = qs_bytes.decode("latin-1") full_url = f"{path}?{qs}" if qs else path query_params = None if qs: parsed = parse_qs(qs) query_params = {k: v[0] if len(v) == 1 else v for k, v in parsed.items()} client = scope.get("client") client_ip = client[0] if client else None return { "endpoint": path, "method": scope.get("method", ""), "path": full_url, "query_params": query_params, "request_body": None, "headers": None, "user_agent": user_agent, "client_ip": client_ip, } async def _log_error( self, request_id: str, request_info: dict, status_code: int, error_type: Optional[str] = None, error_message: Optional[str] = None, error_detail: Optional[dict] = None, stack_trace: Optional[str] = None, response_time_ms: Optional[float] = None, ) -> None: """Log error to database.""" endpoint = request_info["endpoint"] method = request_info["method"] # Don't log errors for DELETE requests to log management endpoints. if method == "DELETE" and ( endpoint.startswith("/api/v1/admin/requests/logs") or endpoint.startswith("/api/v1/admin/errors/logs") ): return try: from app.core.database import AsyncSessionLocal async with AsyncSessionLocal() as db: if error_detail and isinstance(error_detail, dict): if not error_type: error_type = error_detail.get("error_type", "UNKNOWN_ERROR") if not error_message: error_message = error_detail.get("message", "Unknown error occurred") if not error_type: error_type = "INTERNAL_SERVER_ERROR" if status_code >= 500 else "CLIENT_ERROR" if not error_message: error_message = f"HTTP {status_code} error" db.add(ErrorLog( request_id=request_id, endpoint=endpoint, method=method, path=request_info["path"], query_params=request_info["query_params"], request_body=request_info["request_body"], headers=request_info["headers"], error_type=error_type, error_message=error_message, error_detail=error_detail, status_code=status_code, stack_trace=stack_trace, user_agent=request_info["user_agent"], client_ip=request_info["client_ip"], response_time_ms=response_time_ms, )) await db.commit() logger.error( f"Error logged - Request ID: {request_id}, " f"Endpoint: {endpoint}, Status: {status_code}, " f"Error: {error_type} - {error_message}" ) except Exception as e: logger.error(f"Failed to log error to database: {e}") logger.error( f"Original error - Request ID: {request_id}, " f"Status: {status_code}, Error: {error_type} - {error_message}" ) async def _log_request( self, request_id: str, request_info: dict, status_code: int, response_time_ms: Optional[float] = None, response_size: Optional[int] = None, data_source: Optional[str] = None, ) -> None: """Enqueue a request log entry for async batch flushing. Non-blocking: the entry is placed in _REQUEST_LOG_QUEUE and written to the database in bulk by the background flusher, avoiding a synchronous DB round-trip on every request. """ endpoint = request_info["endpoint"] method = request_info["method"] if method == "DELETE" and ( endpoint.startswith("/api/v1/admin/requests/logs") or endpoint.startswith("/api/v1/admin/errors/logs") ): return headers: dict = {} if data_source: headers["X-Data-Source"] = data_source entry = { "request_id": request_id, "endpoint": endpoint, "method": method, "path": request_info["path"], "query_params": request_info["query_params"], "request_body": request_info["request_body"], "headers": headers or None, "status_code": status_code, "response_size": response_size, "user_agent": request_info["user_agent"], "client_ip": request_info["client_ip"], "response_time_ms": response_time_ms, } try: _REQUEST_LOG_QUEUE.put_nowait(entry) except asyncio.QueueFull: logger.warning("Request log queue full — writing synchronously") try: from app.core.database import AsyncSessionLocal async with AsyncSessionLocal() as db: db.add(RequestLog(**entry)) await db.commit() except Exception as e: logger.error(f"Failed to write request log (fallback): {e}")