Coverage for src/lexigram/web/routing/debug.py: 0%

101 statements  

« prev     ^ index     » next       coverage.py v7.15.4, created at 2026-08-25 04:37 +0800

1from __future__ import annotations 

2 

3import ipaddress 

4from typing import TYPE_CHECKING, Any, cast 

5 

6from starlette.responses import JSONResponse 

7 

8from lexigram.logging import get_logger 

9 

10if TYPE_CHECKING: 

11 from lexigram.web.di.provider import WebProvider 

12 

13logger = get_logger(__name__) 

14 

15 

16async def _check_debug_auth( 

17 request: Any, 

18 provider: WebProvider, 

19) -> JSONResponse | None: 

20 """Check debug route authorization. 

21 

22 Returns None when the request is authorized, or a denial JSONResponse 

23 (403 / 404) when it is not. 

24 """ 

25 cfg = provider.web_config 

26 authorized = False 

27 

28 # Token check 

29 token = getattr(cfg, "debug_routes_token", None) or getattr( 

30 provider.provider_config, "debug_routes_token", None 

31 ) 

32 token_value = ( 

33 token.get_secret_value() 

34 if token is not None and hasattr(token, "get_secret_value") 

35 else token 

36 ) 

37 if token_value and request.headers.get("X-Debug-Token") == token_value: 

38 authorized = True 

39 

40 # Custom auth callback 

41 if not authorized and provider.debug_routes_auth: 

42 authorized = await provider.debug_routes_auth(request) 

43 

44 # Middleware gate 

45 req_mw = getattr(cfg, "debug_routes_require_middleware", None) or getattr( 

46 provider.provider_config, "debug_routes_require_middleware", None 

47 ) 

48 if req_mw: 

49 middleware_names = [type(mw).__name__ for mw in provider.middleware] 

50 if req_mw in middleware_names: 

51 authorized = True 

52 else: 

53 return JSONResponse({"error": "Not Found"}, status_code=404) 

54 

55 # IP fallback — allow local, deny remote when no other auth is configured 

56 if not authorized and not token and not provider.debug_routes_auth and not req_mw: 

57 client_ip = _get_client_ip(request) 

58 if not client_ip or _is_local_ip(client_ip): 

59 authorized = True 

60 

61 if not authorized: 

62 return JSONResponse({"error": "Forbidden"}, status_code=403) 

63 

64 return None 

65 

66 

67async def _apply_rate_limit( 

68 request: Any, 

69 provider: WebProvider, 

70) -> JSONResponse | None: 

71 """Apply rate limiting to a debug route request. 

72 

73 Returns None when the request is within limits, or a 429 JSONResponse 

74 when the rate limit is exceeded. 

75 """ 

76 cfg = provider.web_config 

77 try: 

78 rate_limit_raw = getattr(cfg, "debug_routes_rate_limit", 0) or getattr( 

79 provider.provider_config, "debug_routes_rate_limit", 0 

80 ) 

81 rate_limit = ( 

82 int(str(rate_limit_raw)) 

83 if rate_limit_raw is not None and not isinstance(rate_limit_raw, bool) 

84 else 0 

85 ) 

86 if isinstance(rate_limit_raw, bool) and rate_limit_raw: 

87 rate_limit = 100 

88 except (ValueError, TypeError): 

89 rate_limit = 0 

90 

91 logger.debug( 

92 "Debug route rate limit: %s (client: %s)", 

93 rate_limit, 

94 provider._debug_redis_client, 

95 ) 

96 

97 if rate_limit > 0 and provider._debug_redis_client: 

98 try: 

99 window_raw = getattr( 

100 cfg, "debug_routes_rate_window_seconds", 60 

101 ) or getattr( 

102 provider.provider_config, "debug_routes_rate_window_seconds", 60 

103 ) 

104 window = int(str(window_raw)) if window_raw is not None else 60 

105 except (ValueError, TypeError): 

106 window = 60 

107 

108 client_ip = _get_client_ip(request) or "unknown" 

109 key = f"debug_rate_limit:{client_ip}" 

110 

111 try: 

112 count = await provider._debug_redis_client.incr(key) 

113 if count == 1: 

114 await provider._debug_redis_client.expire(key, window) 

115 if count > rate_limit: 

116 return JSONResponse({"error": "Rate limit exceeded"}, status_code=429) 

117 except Exception as e: # noqa: BLE001 

118 logger.warning("Debug rate limiting error: %s", e) 

119 

120 return None 

121 

122 

123def register_debug_routes(provider: WebProvider) -> None: 

124 """Register the /debug/routes diagnostic endpoint.""" 

125 starlette = provider.starlette 

126 if starlette is None: 

127 raise RuntimeError("Starlette application not initialized in provider") 

128 

129 async def debug_routes(request: Any) -> JSONResponse: 

130 deny = await _check_debug_auth(request, provider) 

131 if deny is not None: 

132 return deny 

133 

134 deny = await _apply_rate_limit(request, provider) 

135 if deny is not None: 

136 return deny 

137 

138 server_cfg = getattr(provider.web_config, "server", None) 

139 redact = not getattr(server_cfg, "debug", False) 

140 diagnostics = [] 

141 

142 if hasattr(provider, "router_manager") and hasattr( 

143 provider.router_manager, "_registered_routes" 

144 ): 

145 for ( 

146 method, 

147 path, 

148 ), origins in provider.router_manager._registered_routes.items(): 

149 sanitized = [_sanitize_origin(o, redact) for o in origins] 

150 diagnostics.append( 

151 {"method": method, "path": path, "origins": sanitized} 

152 ) 

153 

154 return JSONResponse({"routes": diagnostics}) 

155 

156 starlette.add_route("/debug/routes", debug_routes, methods=["GET"]) 

157 

158 

159def _get_client_ip(request: Any) -> str | None: 

160 """Get client IP from request headers or connection info.""" 

161 xff = request.headers.get("X-Forwarded-For") 

162 if xff: 

163 return cast("str", xff.split(",")[0].strip()) 

164 

165 xri = request.headers.get("X-Real-IP") 

166 if xri: 

167 return cast("str", xri.strip()) 

168 

169 if hasattr(request, "client") and request.client: 

170 return cast("str", request.client.host) 

171 

172 return None 

173 

174 

175def _sanitize_origin(origin: Any, redact: bool) -> Any: 

176 """Sanitize origin information for debug output.""" 

177 if not isinstance(origin, dict): 

178 return origin 

179 

180 sanitized = origin.copy() 

181 

182 if redact: 

183 if "registered_file" in sanitized: 

184 sanitized["registered_file"] = "<REDACTED>" 

185 if sanitized.get("registered_stack"): 

186 sanitized["registered_stack"] = ["<REDACTED>"] * len( 

187 sanitized["registered_stack"] 

188 ) 

189 

190 return sanitized 

191 

192 

193def _is_local_ip(ip: str) -> bool: 

194 """Check if IP address is local/private.""" 

195 try: 

196 addr = ipaddress.ip_address(ip) 

197 return addr.is_loopback or addr.is_private or addr.is_link_local 

198 except ValueError: 

199 return True