You cannot select more than 25 topics Topics must start with a letter or number, can include dashes ('-') and can be up to 35 characters long.

407 lines
16 KiB
Python

"""
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 100 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.
"""
import asyncio
import json
import time
import traceback
import uuid
from datetime import datetime, timezone
from typing import Callable, Optional
from fastapi import Request, Response
from fastapi.responses import JSONResponse
from sqlalchemy.ext.asyncio import AsyncSession
from starlette.middleware.base import BaseHTTPMiddleware
from starlette.types import ASGIApp
from app.core.database import get_db
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
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)
def start_request_log_flusher() -> None:
"""Schedule the background flusher coroutine. Call once at app startup."""
asyncio.ensure_future(_request_log_flusher())
logger.info("Request log flusher background task started")
class ErrorLoggingMiddleware(BaseHTTPMiddleware):
"""Middleware to log all API requests and errors to database"""
def __init__(self, app: ASGIApp):
super().__init__(app)
async def dispatch(self, request: Request, call_next: Callable) -> Response:
"""Process request and log any errors that occur"""
# Generate unique request ID
request_id = str(uuid.uuid4())[:8]
request.state.request_id = request_id
# Track request start time
start_time = time.time()
# Store request details for potential error logging
request_info = self._extract_request_info_minimal(request)
logger.info(f"Processing request {request_id}: {request.method} {request.url.path}")
try:
# Process the request
response = await call_next(request)
# Calculate response time
response_time_ms = (time.time() - start_time) * 1000
# Store response body for error cases
response_body = b""
error_detail = None
# Check if response indicates an error (4xx or 5xx)
if response.status_code >= 400:
logger.info(f"Error response detected: {response.status_code} for request {request_id}")
# Full extraction for error logging (body + headers)
request_info = await self._extract_request_info(request)
# Try to capture response body for errors
# We need to consume the response body and recreate it
from starlette.responses import Response
# Collect response body chunks
body_chunks = []
async for chunk in response.body_iterator:
body_chunks.append(chunk)
response_body = b''.join(body_chunks)
# Try to parse as JSON
try:
if response_body:
error_detail = json.loads(response_body.decode('utf-8'))
except Exception as e:
logger.warning(f"Could not parse error response body as JSON: {e}")
# Store raw text if not JSON
try:
error_detail = {"raw_response": response_body.decode('utf-8')}
except:
error_detail = {"raw_response": str(response_body)}
# Log the error with response body
logger.info(f"Logging error for request {request_id}")
await self._log_error(
request_id=request_id,
request_info=request_info,
status_code=response.status_code,
error_detail=error_detail,
response_time_ms=response_time_ms
)
# Recreate response with the same body
response = Response(
content=response_body,
status_code=response.status_code,
headers=dict(response.headers),
media_type=response.media_type
)
# Log all requests (not just errors)
# Add data source info to headers for successful responses
data_source = response.headers.get("X-Data-Source", None)
await self._log_request(
request_id=request_id,
request_info=request_info,
status_code=response.status_code,
response_time_ms=response_time_ms,
response_size=len(response_body) if response_body else None,
data_source=data_source
)
# Add request ID to response headers
response.headers["X-Request-ID"] = request_id
return response
except Exception as e:
# Log unexpected errors
response_time_ms = (time.time() - start_time) * 1000
# Full extraction for error logging
request_info = await self._extract_request_info(request)
await self._log_error(
request_id=request_id,
request_info=request_info,
status_code=500,
error_type="INTERNAL_SERVER_ERROR",
error_message=str(e),
stack_trace=traceback.format_exc(),
response_time_ms=response_time_ms
)
# Return error response
return 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}
)
def _extract_request_info_minimal(self, request: Request) -> dict:
"""Extract minimal request info (sync, no body read) for success-path logging."""
return {
"endpoint": str(request.url.path),
"method": request.method,
"path": str(request.url),
"query_params": dict(request.query_params) if request.query_params else None,
"request_body": None,
"headers": None,
"user_agent": request.headers.get("user-agent"),
"client_ip": request.client.host if request.client else None,
}
async def _extract_request_info(self, request: Request) -> dict:
"""Extract full request information for error logging"""
# Get request body if present
body = None
if request.method in ["POST", "PUT", "PATCH"]:
try:
body_bytes = await request.body()
if body_bytes:
body = json.loads(body_bytes.decode('utf-8'))
# Store body for later use in request processing
request._body = body_bytes
except Exception as e:
logger.warning(f"Could not parse request body: {e}")
# Extract headers (remove sensitive ones)
headers = dict(request.headers)
sensitive_headers = ['authorization', 'api-key', 'x-api-key', 'cookie']
for header in sensitive_headers:
if header in headers:
headers[header] = '***REDACTED***'
return {
"endpoint": str(request.url.path),
"method": request.method,
"path": str(request.url),
"query_params": dict(request.query_params) if request.query_params else None,
"request_body": body,
"headers": headers,
"user_agent": headers.get("user-agent"),
"client_ip": request.client.host if request.client else None
}
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
):
"""Log error to database"""
logger.info(f"_log_error called for request {request_id}, status {status_code}")
# Skip logging errors for log deletion endpoints to avoid logging the deletion of logs
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"))):
logger.info(f"Skipping error log for log deletion endpoint: {method} {endpoint}")
return
try:
# Get database session
from app.core.database import AsyncSessionLocal
logger.info(f"Creating database session for request {request_id}")
async with AsyncSessionLocal() as db:
# Extract error information from detail if available
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")
# Set defaults
if not error_type:
if status_code >= 500:
error_type = "INTERNAL_SERVER_ERROR"
elif status_code >= 400:
error_type = "CLIENT_ERROR"
else:
error_type = "UNKNOWN_ERROR"
if not error_message:
error_message = f"HTTP {status_code} error"
# Create error log entry
error_log = ErrorLog(
request_id=request_id,
endpoint=request_info["endpoint"],
method=request_info["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
)
db.add(error_log)
logger.info(f"Error log added to session for request {request_id}")
await db.commit()
logger.info(f"Error log committed to database for request {request_id}")
logger.error(
f"Error logged - Request ID: {request_id}, "
f"Endpoint: {request_info['endpoint']}, "
f"Status: {status_code}, "
f"Error: {error_type} - {error_message}"
)
except Exception as e:
# If we can't log to database, at least log to file
logger.error(f"Failed to log error to database: {e}")
logger.error(
f"Original error - Request ID: {request_id}, "
f"Status: {status_code}, "
f"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
):
"""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.
"""
# Skip logging for log deletion endpoints
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
# Augment headers with data source if provided
headers = request_info.get("headers") or {}
if data_source:
headers = {**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,
"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:
# Queue is at capacity — fall back to a direct write so we don't drop the log
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}")