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
« 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
6from fastapi import APIRouter as _APIRouter
7from fastapi.routing import APIRoute
8from sqlalchemy.ext.asyncio import AsyncSession
10from polar.config import settings
11from polar.kit.pagination import ListResource
12from polar.openapi import APITag
15class AutoCommitAPIRoute(APIRoute):
16 """
17 A subclass of `APIRoute` that automatically
18 commits the session after the endpoint is called.
20 It allows to directly return ORM objects from the endpoint
21 without having to call `session.commit()` before returning.
22 """
24 def __init__(self, path: str, endpoint: Callable[..., Any], **kwargs: Any) -> None:
25 endpoint = self.wrap_endpoint(endpoint)
26 super().__init__(path, endpoint, **kwargs)
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
37 response = await endpoint(*args, **kwargs)
39 if session is not None:
40 await session.commit()
42 return response
44 return wrapped_endpoint
47class IncludedInSchemaAPIRoute(APIRoute):
48 """
49 A subclass of `APIRoute` that automatically sets the `include_in_schema` property
50 depending on the tags.
51 """
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
65class SpeakeasyNameOverrideAPIRoute(APIRoute):
66 """
67 A subclass of `APIRoute` that automatically adds `x-speakeasy-name-override` property
68 following the route function name.
69 """
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 }
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 """
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}
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 """
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 }
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 """
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 }
152class SpeakeasyMCPAPIRoute(APIRoute):
153 """
154 A subclass of `APIRoute` that automatically adds `x-speakeasy-mcp` property
155 to the OpenAPI schema.
156 """
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))
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
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
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 """
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)
198 return _CustomAPIRouter
201__all__ = [
202 "get_api_router_class",
203 "AutoCommitAPIRoute",
204 "IncludedInSchemaAPIRoute",
205 "SpeakeasyGroupAPIRoute",
206 "SpeakeasyIgnoreAPIRoute",
207 "SpeakeasyNameOverrideAPIRoute",
208 "SpeakeasyPaginationAPIRoute",
209 "SpeakeasyMCPAPIRoute",
210]