Coverage for agentos/api/middleware.py: 57%
74 statements
« prev ^ index » next coverage.py v7.14.3, created at 2026-07-09 09:19 +0800
« prev ^ index » next coverage.py v7.14.3, created at 2026-07-09 09:19 +0800
1"""AgentOS API middleware — request/response processing pipeline.
3Provides authentication, CORS, request tracing, request-ID injection,
4and rate-limiting middleware for the AgentOS API server.
5"""
7from __future__ import annotations
9import time
10import uuid
11from collections.abc import Callable
12from dataclasses import dataclass, field
14# ── Request ID ────────────────────────────────────────────────────────────────
17@dataclass
18class RequestContext:
19 """请求上下文。"""
21 request_id: str = ""
22 start_time: float = 0.0
23 method: str = ""
24 path: str = ""
25 client_ip: str = ""
26 user_agent: str = ""
28 @property
29 def elapsed_ms(self) -> float:
30 return (time.monotonic() - self.start_time) * 1000
33class RequestIDMiddleware:
34 """Inject X-Request-ID into every request and propagate it."""
36 def __init__(self, header: str = "X-Request-ID"):
37 self.header = header
39 def process_request(self, headers: dict) -> RequestContext:
40 rid = headers.get(self.header.lower(), headers.get(self.header, ""))
41 if not rid:
42 rid = str(uuid.uuid4())[:12]
43 return RequestContext(request_id=rid, start_time=time.monotonic())
46# ── CORS ──────────────────────────────────────────────────────────────────────
49@dataclass
50class CORSConfig:
51 """CORS 配置。"""
53 allow_origins: list[str] = field(default_factory=lambda: ["*"])
54 allow_methods: list[str] = field(
55 default_factory=lambda: ["GET", "POST", "PUT", "DELETE", "PATCH", "OPTIONS"]
56 )
57 allow_headers: list[str] = field(
58 default_factory=lambda: ["Content-Type", "Authorization", "X-Request-ID"]
59 )
60 expose_headers: list[str] = field(default_factory=list)
61 max_age: int = 86400
62 allow_credentials: bool = False
65class CORSMiddleware:
66 """Add CORS headers to every response."""
68 def __init__(self, config: CORSConfig | None = None):
69 self.config = config or CORSConfig()
71 def apply(self, response_headers: dict) -> dict:
72 origin = self.config.allow_origins[0] if self.config.allow_origins else ""
73 response_headers["Access-Control-Allow-Origin"] = origin
74 response_headers["Access-Control-Allow-Methods"] = ", ".join(self.config.allow_methods)
75 response_headers["Access-Control-Allow-Headers"] = ", ".join(self.config.allow_headers)
76 if self.config.expose_headers:
77 response_headers["Access-Control-Expose-Headers"] = ", ".join(
78 self.config.expose_headers
79 )
80 response_headers["Access-Control-Max-Age"] = str(self.config.max_age)
81 if self.config.allow_credentials:
82 response_headers["Access-Control-Allow-Credentials"] = "true"
83 return response_headers
86# ── Auth ──────────────────────────────────────────────────────────────────────
89@dataclass
90class AuthConfig:
91 """认证配置。"""
93 api_key_header: str = "X-API-Key"
94 api_key: str = ""
95 enabled: bool = True
98class AuthMiddleware:
99 """Simple API-key authentication middleware."""
101 def __init__(self, config: AuthConfig | None = None):
102 self.config = config or AuthConfig()
104 def authenticate(self, headers: dict) -> tuple[bool, str]:
105 """Return (authorized, message)."""
106 if not self.config.enabled or not self.config.api_key:
107 return True, ""
108 provided = headers.get(
109 self.config.api_key_header.lower(), headers.get(self.config.api_key_header, "")
110 )
111 if provided != self.config.api_key:
112 return False, "Invalid or missing API key"
113 return True, ""
116# ── Request logger ────────────────────────────────────────────────────────────
119class RequestLogMiddleware:
120 """Log every request with method, path, status, and elapsed time."""
122 def __init__(self, logger: Callable[[str], None] | None = None):
123 self._log = logger or print
125 def log(self, ctx: RequestContext, status: int) -> str:
126 msg = f"[{ctx.request_id}] {ctx.method} {ctx.path} " f"→ {status} ({ctx.elapsed_ms:.1f}ms)"
127 self._log(msg)
128 return msg
131# ── Middleware stack ──────────────────────────────────────────────────────────
134class MiddlewareStack:
135 """Ordered middleware pipeline for the AgentOS API."""
137 def __init__(
138 self,
139 cors: CORSMiddleware | None = None,
140 auth: AuthMiddleware | None = None,
141 req_log: RequestLogMiddleware | None = None,
142 req_id: RequestIDMiddleware | None = None,
143 ):
144 self.cors = cors or CORSMiddleware()
145 self.auth = auth or AuthMiddleware()
146 self.req_log = req_log or RequestLogMiddleware()
147 self.req_id = req_id or RequestIDMiddleware()