Coverage for polar/kit/routing.py: 96%

88 statements  

« prev     ^ index     » next       coverage.py v7.15.2, created at 2026-10-07 12:42 +0000

1import functools 

2import inspect 

3from collections.abc import Callable 

4from typing import Any 

5 

6from fastapi import APIRouter as _APIRouter 

7from fastapi.routing import APIRoute 

8from sqlalchemy.ext.asyncio import AsyncSession 

9 

10from polar.config import settings 

11from polar.kit.pagination import ListResource 

12from polar.openapi import APITag 

13 

14 

15class AutoCommitAPIRoute(APIRoute): 

16 """ 

17 A subclass of `APIRoute` that automatically 

18 commits the session after the endpoint is called. 

19 

20 It allows to directly return ORM objects from the endpoint 

21 without having to call `session.commit()` before returning. 

22 """ 

23 

24 def __init__(self, path: str, endpoint: Callable[..., Any], **kwargs: Any) -> None: 

25 endpoint = self.wrap_endpoint(endpoint) 

26 super().__init__(path, endpoint, **kwargs) 

27 

28 def wrap_endpoint(self, endpoint: Callable[..., Any]) -> Callable[..., Any]: 

29 @functools.wraps(endpoint) 

30 async def wrapped_endpoint(*args: Any, **kwargs: Any) -> Any: 

31 session: AsyncSession | None = None 

32 for arg in (args, *kwargs.values()): 

33 if isinstance(arg, AsyncSession): 

34 session = arg 

35 break 

36 

37 response = await endpoint(*args, **kwargs) 

38 

39 if session is not None: 

40 await session.commit() 

41 

42 return response 

43 

44 return wrapped_endpoint 

45 

46 

47class IncludedInSchemaAPIRoute(APIRoute): 

48 """ 

49 A subclass of `APIRoute` that automatically sets the `include_in_schema` property 

50 depending on the tags. 

51 """ 

52 

53 def __init__(self, path: str, endpoint: Callable[..., Any], **kwargs: Any) -> None: 

54 super().__init__(path, endpoint, **kwargs) 

55 tags = self.tags 

56 if self.include_in_schema: 

57 if APITag.private in tags: 

58 self.include_in_schema = settings.is_development() 

59 elif APITag.public in tags: 59 ↛ 62line 59 didn't jump to line 62 because the condition on line 59 was always true

60 self.include_in_schema = True 

61 else: 

62 self.include_in_schema = False 

63 

64 

65class SpeakeasyNameOverrideAPIRoute(APIRoute): 

66 """ 

67 A subclass of `APIRoute` that automatically adds `x-speakeasy-name-override` property 

68 following the route function name. 

69 """ 

70 

71 def __init__(self, path: str, endpoint: Callable[..., Any], **kwargs: Any) -> None: 

72 super().__init__(path, endpoint, **kwargs) 

73 endpoint_name = endpoint.__name__ 

74 openapi_extra = self.openapi_extra or {} 

75 if "x-speakeasy-name-override" not in openapi_extra: 

76 self.openapi_extra = { 

77 **openapi_extra, 

78 "x-speakeasy-name-override": endpoint_name, 

79 } 

80 

81 

82class SpeakeasyIgnoreAPIRoute(APIRoute): 

83 """ 

84 A subclass of `APIRoute` that automatically adds `x-speakeasy-ignore` property 

85 to the OpenAPI schema if `APITag.documented` is missing. 

86 """ 

87 

88 def __init__(self, path: str, endpoint: Callable[..., Any], **kwargs: Any) -> None: 

89 super().__init__(path, endpoint, **kwargs) 

90 tags = self.tags 

91 if APITag.public not in tags: 

92 openapi_extra = self.openapi_extra or {} 

93 self.openapi_extra = {**openapi_extra, "x-speakeasy-ignore": True} 

94 

95 

96class SpeakeasyGroupAPIRoute(APIRoute): 

97 """ 

98 A subclass of `APIRoute` that automatically adds `x-speakeasy-group` property 

99 to the OpenAPI schema by combining all the non-generic tags. 

100 """ 

101 

102 def __init__(self, path: str, endpoint: Callable[..., Any], **kwargs: Any) -> None: 

103 super().__init__(path, endpoint, **kwargs) 

104 non_generic_tags = [str(tag) for tag in self.tags if tag not in APITag] 

105 if len(non_generic_tags) > 0: 105 ↛ exitline 105 didn't return from function '__init__' because the condition on line 105 was always true

106 openapi_extra = self.openapi_extra or {} 

107 self.openapi_extra = { 

108 **openapi_extra, 

109 "x-speakeasy-group": ".".join(non_generic_tags), 

110 } 

111 

112 

113class SpeakeasyPaginationAPIRoute(APIRoute): 

114 """ 

115 A subclass of `APIRoute` that automatically adds `x-speakeasy-pagination` property 

116 to the OpenAPI schema if the endpoint response model is a `ListResource`. 

117 """ 

118 

119 def __init__(self, path: str, endpoint: Callable[..., Any], **kwargs: Any) -> None: 

120 super().__init__(path, endpoint, **kwargs) 

121 response_model = self.response_model 

122 if ( 

123 response_model is not None 

124 and inspect.isclass(response_model) 

125 and ListResource in response_model.mro() 

126 ): 

127 openapi_extra = self.openapi_extra or {} 

128 self.openapi_extra = { 

129 **openapi_extra, 

130 "x-speakeasy-pagination": { 

131 "type": "offsetLimit", 

132 "inputs": [ 

133 { 

134 "name": "page", 

135 "in": "parameters", 

136 "type": "page", 

137 }, 

138 { 

139 "name": "limit", 

140 "in": "parameters", 

141 "type": "limit", 

142 }, 

143 ], 

144 "outputs": { 

145 "results": "$.items", 

146 "numPages": "$.pagination.max_page", 

147 }, 

148 }, 

149 } 

150 

151 

152class SpeakeasyMCPAPIRoute(APIRoute): 

153 """ 

154 A subclass of `APIRoute` that automatically adds `x-speakeasy-mcp` property 

155 to the OpenAPI schema. 

156 """ 

157 

158 def __init__(self, path: str, endpoint: Callable[..., Any], **kwargs: Any) -> None: 

159 super().__init__(path, endpoint, **kwargs) 

160 openapi_extra = self.openapi_extra or {} 

161 if APITag.mcp in self.tags: 

162 safe_method = all( 

163 method in {"GET", "HEAD", "OPTIONS"} for method in self.methods 

164 ) 

165 scopes = [ 

166 "read" if safe_method else "write", 

167 ] 

168 non_generic_tags = [str(tag) for tag in self.tags if tag not in APITag] 

169 if len(non_generic_tags) > 0: 169 ↛ 172line 169 didn't jump to line 172 because the condition on line 169 was always true

170 scopes.append(".".join(non_generic_tags)) 

171 

172 openapi_extra = { 

173 **openapi_extra, 

174 "x-speakeasy-mcp": {"disabled": False, "scopes": scopes}, 

175 } 

176 else: 

177 openapi_extra = {**openapi_extra, "x-speakeasy-mcp": {"disabled": True}} 

178 self.openapi_extra = openapi_extra 

179 

180 

181def _inherit_signature_from[**P, T]( 

182 _to: Callable[P, T], 

183) -> Callable[[Callable[..., T]], Callable[P, T]]: 

184 return lambda x: x # pyright: ignore 

185 

186 

187def get_api_router_class(route_class: type[APIRoute]) -> type[_APIRouter]: 

188 """ 

189 Returns a subclass of `APIRouter` that uses the given `route_class`. 

190 """ 

191 

192 class _CustomAPIRouter(_APIRouter): 

193 @_inherit_signature_from(_APIRouter.__init__) 

194 def __init__(self, *args: Any, **kwargs: Any) -> None: 

195 kwargs["route_class"] = route_class 

196 super().__init__(*args, **kwargs) 

197 

198 return _CustomAPIRouter 

199 

200 

201__all__ = [ 

202 "get_api_router_class", 

203 "AutoCommitAPIRoute", 

204 "IncludedInSchemaAPIRoute", 

205 "SpeakeasyGroupAPIRoute", 

206 "SpeakeasyIgnoreAPIRoute", 

207 "SpeakeasyNameOverrideAPIRoute", 

208 "SpeakeasyPaginationAPIRoute", 

209 "SpeakeasyMCPAPIRoute", 

210]