Coverage for agentos/api/middleware.py: 57%

74 statements  

« prev     ^ index     » next       coverage.py v7.14.3, created at 2026-07-09 10:19 +0800

1"""AgentOS API middleware — request/response processing pipeline. 

2 

3Provides authentication, CORS, request tracing, request-ID injection, 

4and rate-limiting middleware for the AgentOS API server. 

5""" 

6 

7from __future__ import annotations 

8 

9import time 

10import uuid 

11from collections.abc import Callable 

12from dataclasses import dataclass, field 

13 

14# ── Request ID ──────────────────────────────────────────────────────────────── 

15 

16 

17@dataclass 

18class RequestContext: 

19 """请求上下文。""" 

20 

21 request_id: str = "" 

22 start_time: float = 0.0 

23 method: str = "" 

24 path: str = "" 

25 client_ip: str = "" 

26 user_agent: str = "" 

27 

28 @property 

29 def elapsed_ms(self) -> float: 

30 return (time.monotonic() - self.start_time) * 1000 

31 

32 

33class RequestIDMiddleware: 

34 """Inject X-Request-ID into every request and propagate it.""" 

35 

36 def __init__(self, header: str = "X-Request-ID"): 

37 self.header = header 

38 

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()) 

44 

45 

46# ── CORS ────────────────────────────────────────────────────────────────────── 

47 

48 

49@dataclass 

50class CORSConfig: 

51 """CORS 配置。""" 

52 

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 

63 

64 

65class CORSMiddleware: 

66 """Add CORS headers to every response.""" 

67 

68 def __init__(self, config: CORSConfig | None = None): 

69 self.config = config or CORSConfig() 

70 

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 

84 

85 

86# ── Auth ────────────────────────────────────────────────────────────────────── 

87 

88 

89@dataclass 

90class AuthConfig: 

91 """认证配置。""" 

92 

93 api_key_header: str = "X-API-Key" 

94 api_key: str = "" 

95 enabled: bool = True 

96 

97 

98class AuthMiddleware: 

99 """Simple API-key authentication middleware.""" 

100 

101 def __init__(self, config: AuthConfig | None = None): 

102 self.config = config or AuthConfig() 

103 

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, "" 

114 

115 

116# ── Request logger ──────────────────────────────────────────────────────────── 

117 

118 

119class RequestLogMiddleware: 

120 """Log every request with method, path, status, and elapsed time.""" 

121 

122 def __init__(self, logger: Callable[[str], None] | None = None): 

123 self._log = logger or print 

124 

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 

129 

130 

131# ── Middleware stack ────────────────────────────────────────────────────────── 

132 

133 

134class MiddlewareStack: 

135 """Ordered middleware pipeline for the AgentOS API.""" 

136 

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()