Coverage for src/lexigram/web/middleware/cors.py: 54%

26 statements  

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

1"""CORS middleware wrapper for lexigram-web security primitives.""" 

2 

3from __future__ import annotations 

4 

5from collections.abc import Sequence 

6from typing import Any 

7 

8from lexigram.logging import get_logger 

9from lexigram.web.security.config import CORSConfig 

10from lexigram.web.security.cors.middleware import CORSMiddleware as WebCORSMiddleware 

11 

12logger = get_logger(__name__) 

13 

14 

15class CORSMiddleware(WebCORSMiddleware): 

16 """CORS middleware with web-compatible interface. 

17 

18 Wraps :class:`lexigram.web.security.cors.middleware.CORSMiddleware` 

19 while still accepting Starlette-style keyword arguments. 

20 """ 

21 

22 def __init__( 

23 self, 

24 app: Any, 

25 config: Any = None, 

26 **kwargs: Any, 

27 ) -> None: 

28 """Initialize CORS middleware with optional config conversion. 

29 

30 Args: 

31 app: ASGI application. 

32 config: CORSConfig instance or duck-typed equivalent. 

33 **kwargs: Additional kwargs as fallback for config values. 

34 """ 

35 super().__init__(app=app, config=self._convert_config(config, kwargs)) 

36 

37 @staticmethod 

38 def _convert_config(config: Any, kwargs: dict[str, Any]) -> CORSConfig: 

39 """Convert various config formats to CORSConfig. 

40 

41 Args: 

42 config: Configuration object (CORSConfig or duck-typed equivalent). 

43 kwargs: Keyword arguments as fallback. 

44 

45 Returns: 

46 CORSConfig instance. 

47 """ 

48 if config is None: 

49 return CORSConfig( 

50 allowed_origins=kwargs.get("allow_origins", ["*"]), 

51 allow_credentials=kwargs.get("allow_credentials", False), 

52 allow_methods=kwargs.get( 

53 "allow_methods", 

54 ["GET", "POST", "PUT", "DELETE", "PATCH"], 

55 ), 

56 allow_headers=kwargs.get("allow_headers", ["*"]), 

57 expose_headers=kwargs.get("expose_headers", []), 

58 max_age=kwargs.get("max_age", 600), 

59 ) 

60 

61 if isinstance(config, CORSConfig): 

62 return config 

63 

64 if hasattr(config, "to_middleware_kwargs"): 

65 cors_kwargs = config.to_middleware_kwargs() 

66 return CORSConfig( 

67 allowed_origins=cors_kwargs.get("allow_origins", ["*"]), 

68 allow_credentials=cors_kwargs.get("allow_credentials", False), 

69 allow_methods=cors_kwargs.get( 

70 "allow_methods", ["GET", "POST", "PUT", "DELETE", "PATCH"] 

71 ), 

72 allow_headers=cors_kwargs.get("allow_headers", ["*"]), 

73 expose_headers=cors_kwargs.get("expose_headers", []), 

74 max_age=cors_kwargs.get("max_age", 600), 

75 ) 

76 

77 return CORSConfig( 

78 allowed_origins=getattr(config, "allow_origins", ["*"]), 

79 allow_credentials=getattr(config, "allow_credentials", False), 

80 allow_methods=getattr( 

81 config, 

82 "allow_methods", 

83 ["GET", "POST", "PUT", "DELETE", "PATCH"], 

84 ), 

85 allow_headers=getattr(config, "allow_headers", ["*"]), 

86 expose_headers=getattr(config, "expose_headers", []), 

87 max_age=getattr(config, "max_age", 600), 

88 ) 

89 

90 

91def create_development_cors() -> CORSConfig: 

92 """Create CORS config for development (permissive). 

93 

94 WARNING: Never use in production! 

95 

96 Returns: 

97 Development CORS config 

98 """ 

99 logger.warning( 

100 "Using development CORS configuration. " 

101 "This allows all origins - DO NOT USE IN PRODUCTION!", 

102 ) 

103 

104 return CORSConfig( 

105 allow_origins=["*"], 

106 allow_credentials=False, # Can't use credentials with wildcard 

107 allow_methods=["GET", "POST", "PUT", "DELETE", "PATCH", "OPTIONS"], 

108 ) 

109 

110 

111def create_production_cors( 

112 allowed_domains: Sequence[str], 

113 allow_credentials: bool = True, 

114) -> CORSConfig: 

115 """Create CORS config for production (strict). 

116 

117 Args: 

118 allowed_domains: Exact domains to allow (e.g., ["https://myapp.com"]) 

119 allow_credentials: Whether to allow credentials 

120 

121 Returns: 

122 Production CORS config 

123 """ 

124 return CORSConfig( 

125 allow_origins=list(allowed_domains), 

126 allow_credentials=allow_credentials, 

127 allow_methods=["GET", "POST", "PUT", "DELETE", "PATCH", "OPTIONS"], 

128 expose_headers=["Content-Length", "Content-Type"], 

129 ) 

130 

131 

132__all__ = ["CORSMiddleware", "create_development_cors", "create_production_cors"]