kiba-core 0.5.3.dev61__tar.gz → 0.5.3.dev63__tar.gz
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- {kiba_core-0.5.3.dev61 → kiba_core-0.5.3.dev63}/CHANGELOG.md +9 -0
- {kiba_core-0.5.3.dev61 → kiba_core-0.5.3.dev63}/PKG-INFO +1 -1
- {kiba_core-0.5.3.dev61 → kiba_core-0.5.3.dev63}/core/api/api_request.py +1 -0
- {kiba_core-0.5.3.dev61 → kiba_core-0.5.3.dev63}/core/api/authorizer.py +48 -6
- {kiba_core-0.5.3.dev61 → kiba_core-0.5.3.dev63}/core/api/json_route.py +1 -0
- kiba_core-0.5.3.dev63/core/api/middleware/origin_ip_middleware.py +36 -0
- kiba_core-0.5.3.dev63/core/api/openapi.py +215 -0
- kiba_core-0.5.3.dev63/core/api/rate_limit.py +110 -0
- kiba_core-0.5.3.dev63/core/api/request_context.py +52 -0
- kiba_core-0.5.3.dev63/core/api/route.py +80 -0
- kiba_core-0.5.3.dev63/core/api/route_auth.py +11 -0
- kiba_core-0.5.3.dev63/core/api/route_metadata.py +52 -0
- {kiba_core-0.5.3.dev61 → kiba_core-0.5.3.dev63}/core/api/streaming_json_route.py +1 -0
- {kiba_core-0.5.3.dev61 → kiba_core-0.5.3.dev63}/kiba_core.egg-info/PKG-INFO +1 -1
- {kiba_core-0.5.3.dev61 → kiba_core-0.5.3.dev63}/kiba_core.egg-info/SOURCES.txt +12 -0
- {kiba_core-0.5.3.dev61 → kiba_core-0.5.3.dev63}/pyproject.toml +1 -1
- {kiba_core-0.5.3.dev61 → kiba_core-0.5.3.dev63}/tests/api/test_authorizer.py +36 -18
- {kiba_core-0.5.3.dev61 → kiba_core-0.5.3.dev63}/tests/api/test_exception_handling_middleware.py +12 -0
- {kiba_core-0.5.3.dev61 → kiba_core-0.5.3.dev63}/tests/api/test_json_route.py +8 -0
- kiba_core-0.5.3.dev63/tests/api/test_openapi.py +200 -0
- kiba_core-0.5.3.dev63/tests/api/test_origin_ip_middleware.py +54 -0
- kiba_core-0.5.3.dev63/tests/api/test_rate_limit.py +162 -0
- kiba_core-0.5.3.dev63/tests/api/test_request_context.py +61 -0
- kiba_core-0.5.3.dev63/tests/api/test_route.py +184 -0
- {kiba_core-0.5.3.dev61 → kiba_core-0.5.3.dev63}/tests/api/test_streaming_json_route.py +7 -0
- {kiba_core-0.5.3.dev61 → kiba_core-0.5.3.dev63}/tests/caching/test_file_cache.py +0 -7
- {kiba_core-0.5.3.dev61 → kiba_core-0.5.3.dev63}/tests/test_database.py +7 -13
- {kiba_core-0.5.3.dev61 → kiba_core-0.5.3.dev63}/.github/pull_request_template.md +0 -0
- {kiba_core-0.5.3.dev61 → kiba_core-0.5.3.dev63}/.github/workflows/deploy.yml +0 -0
- {kiba_core-0.5.3.dev61 → kiba_core-0.5.3.dev63}/.github/workflows/pull-request.yml +0 -0
- {kiba_core-0.5.3.dev61 → kiba_core-0.5.3.dev63}/.github/workflows/release.yml +0 -0
- {kiba_core-0.5.3.dev61 → kiba_core-0.5.3.dev63}/.gitignore +0 -0
- {kiba_core-0.5.3.dev61 → kiba_core-0.5.3.dev63}/Dockerfile +0 -0
- {kiba_core-0.5.3.dev61 → kiba_core-0.5.3.dev63}/README.md +0 -0
- {kiba_core-0.5.3.dev61 → kiba_core-0.5.3.dev63}/core/__init__.py +0 -0
- {kiba_core-0.5.3.dev61 → kiba_core-0.5.3.dev63}/core/api/__init__.py +0 -0
- {kiba_core-0.5.3.dev61 → kiba_core-0.5.3.dev63}/core/api/api_response.py +0 -0
- {kiba_core-0.5.3.dev61 → kiba_core-0.5.3.dev63}/core/api/default_routes.py +0 -0
- {kiba_core-0.5.3.dev61 → kiba_core-0.5.3.dev63}/core/api/health.py +0 -0
- {kiba_core-0.5.3.dev61 → kiba_core-0.5.3.dev63}/core/api/middleware/__init__.py +0 -0
- {kiba_core-0.5.3.dev61 → kiba_core-0.5.3.dev63}/core/api/middleware/database_connection_middleware.py +0 -0
- {kiba_core-0.5.3.dev61 → kiba_core-0.5.3.dev63}/core/api/middleware/exception_handling_middleware.py +0 -0
- {kiba_core-0.5.3.dev61 → kiba_core-0.5.3.dev63}/core/api/middleware/logging_middleware.py +0 -0
- {kiba_core-0.5.3.dev61 → kiba_core-0.5.3.dev63}/core/api/middleware/server_headers_middleware.py +0 -0
- {kiba_core-0.5.3.dev61 → kiba_core-0.5.3.dev63}/core/aws_requester.py +0 -0
- {kiba_core-0.5.3.dev61 → kiba_core-0.5.3.dev63}/core/caching/__init__.py +0 -0
- {kiba_core-0.5.3.dev61 → kiba_core-0.5.3.dev63}/core/caching/cache.py +0 -0
- {kiba_core-0.5.3.dev61 → kiba_core-0.5.3.dev63}/core/caching/dict_cache.py +0 -0
- {kiba_core-0.5.3.dev61 → kiba_core-0.5.3.dev63}/core/caching/file_cache.py +0 -0
- {kiba_core-0.5.3.dev61 → kiba_core-0.5.3.dev63}/core/exceptions.py +0 -0
- {kiba_core-0.5.3.dev61 → kiba_core-0.5.3.dev63}/core/http/__init__.py +0 -0
- {kiba_core-0.5.3.dev61 → kiba_core-0.5.3.dev63}/core/http/basic_authentication.py +0 -0
- {kiba_core-0.5.3.dev61 → kiba_core-0.5.3.dev63}/core/http/jwt.py +0 -0
- {kiba_core-0.5.3.dev61 → kiba_core-0.5.3.dev63}/core/http/rest_method.py +0 -0
- {kiba_core-0.5.3.dev61 → kiba_core-0.5.3.dev63}/core/logging.py +0 -0
- {kiba_core-0.5.3.dev61 → kiba_core-0.5.3.dev63}/core/notifications/__init__.py +0 -0
- {kiba_core-0.5.3.dev61 → kiba_core-0.5.3.dev63}/core/notifications/discord_client.py +0 -0
- {kiba_core-0.5.3.dev61 → kiba_core-0.5.3.dev63}/core/notifications/notification_client.py +0 -0
- {kiba_core-0.5.3.dev61 → kiba_core-0.5.3.dev63}/core/notifications/slack_client.py +0 -0
- {kiba_core-0.5.3.dev61 → kiba_core-0.5.3.dev63}/core/py.typed +0 -0
- {kiba_core-0.5.3.dev61 → kiba_core-0.5.3.dev63}/core/queues/__init__.py +0 -0
- {kiba_core-0.5.3.dev61 → kiba_core-0.5.3.dev63}/core/queues/aqs.py +0 -0
- {kiba_core-0.5.3.dev61 → kiba_core-0.5.3.dev63}/core/queues/cosmos.py +0 -0
- {kiba_core-0.5.3.dev61 → kiba_core-0.5.3.dev63}/core/queues/message_queue.py +0 -0
- {kiba_core-0.5.3.dev61 → kiba_core-0.5.3.dev63}/core/queues/message_queue_processor.py +0 -0
- {kiba_core-0.5.3.dev61 → kiba_core-0.5.3.dev63}/core/queues/model.py +0 -0
- {kiba_core-0.5.3.dev61 → kiba_core-0.5.3.dev63}/core/queues/sql.py +0 -0
- {kiba_core-0.5.3.dev61 → kiba_core-0.5.3.dev63}/core/queues/sqs.py +0 -0
- {kiba_core-0.5.3.dev61 → kiba_core-0.5.3.dev63}/core/requester/__init__.py +0 -0
- {kiba_core-0.5.3.dev61 → kiba_core-0.5.3.dev63}/core/requester/requester.py +0 -0
- {kiba_core-0.5.3.dev61 → kiba_core-0.5.3.dev63}/core/s3_manager.py +0 -0
- {kiba_core-0.5.3.dev61 → kiba_core-0.5.3.dev63}/core/service_client.py +0 -0
- {kiba_core-0.5.3.dev61 → kiba_core-0.5.3.dev63}/core/store/__init__.py +0 -0
- {kiba_core-0.5.3.dev61 → kiba_core-0.5.3.dev63}/core/store/database.py +0 -0
- {kiba_core-0.5.3.dev61 → kiba_core-0.5.3.dev63}/core/store/retriever.py +0 -0
- {kiba_core-0.5.3.dev61 → kiba_core-0.5.3.dev63}/core/store/saver.py +0 -0
- {kiba_core-0.5.3.dev61 → kiba_core-0.5.3.dev63}/core/util/__init__.py +0 -0
- {kiba_core-0.5.3.dev61 → kiba_core-0.5.3.dev63}/core/util/async_util.py +0 -0
- {kiba_core-0.5.3.dev61 → kiba_core-0.5.3.dev63}/core/util/chain_util.py +0 -0
- {kiba_core-0.5.3.dev61 → kiba_core-0.5.3.dev63}/core/util/date_util.py +0 -0
- {kiba_core-0.5.3.dev61 → kiba_core-0.5.3.dev63}/core/util/dict_util.py +0 -0
- {kiba_core-0.5.3.dev61 → kiba_core-0.5.3.dev63}/core/util/file_util.py +0 -0
- {kiba_core-0.5.3.dev61 → kiba_core-0.5.3.dev63}/core/util/hashing_util.py +0 -0
- {kiba_core-0.5.3.dev61 → kiba_core-0.5.3.dev63}/core/util/http_util.py +0 -0
- {kiba_core-0.5.3.dev61 → kiba_core-0.5.3.dev63}/core/util/json_util.py +0 -0
- {kiba_core-0.5.3.dev61 → kiba_core-0.5.3.dev63}/core/util/list_util.py +0 -0
- {kiba_core-0.5.3.dev61 → kiba_core-0.5.3.dev63}/core/util/string_util.py +0 -0
- {kiba_core-0.5.3.dev61 → kiba_core-0.5.3.dev63}/core/util/typing_util.py +0 -0
- {kiba_core-0.5.3.dev61 → kiba_core-0.5.3.dev63}/core/util/url_util.py +0 -0
- {kiba_core-0.5.3.dev61 → kiba_core-0.5.3.dev63}/core/util/value_holder.py +0 -0
- {kiba_core-0.5.3.dev61 → kiba_core-0.5.3.dev63}/core/web3/__init__.py +0 -0
- {kiba_core-0.5.3.dev61 → kiba_core-0.5.3.dev63}/core/web3/eth_client.py +0 -0
- {kiba_core-0.5.3.dev61 → kiba_core-0.5.3.dev63}/core/web3/multicall3.py +0 -0
- {kiba_core-0.5.3.dev61 → kiba_core-0.5.3.dev63}/kiba_core.egg-info/dependency_links.txt +0 -0
- {kiba_core-0.5.3.dev61 → kiba_core-0.5.3.dev63}/kiba_core.egg-info/requires.txt +0 -0
- {kiba_core-0.5.3.dev61 → kiba_core-0.5.3.dev63}/kiba_core.egg-info/top_level.txt +0 -0
- {kiba_core-0.5.3.dev61 → kiba_core-0.5.3.dev63}/makefile +0 -0
- {kiba_core-0.5.3.dev61 → kiba_core-0.5.3.dev63}/setup.cfg +0 -0
- {kiba_core-0.5.3.dev61 → kiba_core-0.5.3.dev63}/tests/__init__.py +0 -0
- {kiba_core-0.5.3.dev61 → kiba_core-0.5.3.dev63}/tests/api/test_api_response.py +0 -0
- {kiba_core-0.5.3.dev61 → kiba_core-0.5.3.dev63}/tests/caching/test_dict_cache.py +0 -0
- {kiba_core-0.5.3.dev61 → kiba_core-0.5.3.dev63}/tests/requester/__init__.py +0 -0
- {kiba_core-0.5.3.dev61 → kiba_core-0.5.3.dev63}/tests/requester/test_requester.py +0 -0
- {kiba_core-0.5.3.dev61 → kiba_core-0.5.3.dev63}/tests/test_exceptions.py +0 -0
- {kiba_core-0.5.3.dev61 → kiba_core-0.5.3.dev63}/tests/util/__init__.py +0 -0
- {kiba_core-0.5.3.dev61 → kiba_core-0.5.3.dev63}/tests/util/test_async_util.py +0 -0
- {kiba_core-0.5.3.dev61 → kiba_core-0.5.3.dev63}/tests/util/test_chain_util.py +0 -0
- {kiba_core-0.5.3.dev61 → kiba_core-0.5.3.dev63}/tests/util/test_date_util.py +0 -0
- {kiba_core-0.5.3.dev61 → kiba_core-0.5.3.dev63}/tests/util/test_dict_util.py +0 -0
- {kiba_core-0.5.3.dev61 → kiba_core-0.5.3.dev63}/tests/util/test_file_util.py +0 -0
- {kiba_core-0.5.3.dev61 → kiba_core-0.5.3.dev63}/tests/util/test_json_util.py +0 -0
- {kiba_core-0.5.3.dev61 → kiba_core-0.5.3.dev63}/tests/util/test_list_util.py +0 -0
- {kiba_core-0.5.3.dev61 → kiba_core-0.5.3.dev63}/tests/util/test_string_util.py +0 -0
- {kiba_core-0.5.3.dev61 → kiba_core-0.5.3.dev63}/tests/util/test_url_util.py +0 -0
- {kiba_core-0.5.3.dev61 → kiba_core-0.5.3.dev63}/tests/web3/__init__.py +0 -0
- {kiba_core-0.5.3.dev61 → kiba_core-0.5.3.dev63}/tests/web3/test_rest_eth_client.py +0 -0
- {kiba_core-0.5.3.dev61 → kiba_core-0.5.3.dev63}/uv.lock +0 -0
|
@@ -8,6 +8,15 @@ The format is based on [Keep a Changelog](https://keepachangelog.com/en/1.0.0/)
|
|
|
8
8
|
## [Unreleased]
|
|
9
9
|
|
|
10
10
|
### Added
|
|
11
|
+
- [MAJOR] Reworked `route` to compose JSON or streaming transport, application-resolved authorization, and rate limiting
|
|
12
|
+
- [MAJOR] Replaced authorization-decorator sequences with named `RouteAuthResolver` policies
|
|
13
|
+
- [MINOR] Added `authJwt`, `authBasic`, and `originIp` to `KibaApiRequest`
|
|
14
|
+
- [MAJOR] Replaced global request-context accessors with injectable `RequestContextHolder` and `RequestContextMiddleware`
|
|
15
|
+
- [MINOR] Added `RateLimitConfig` for rate-limit windows and request keying
|
|
16
|
+
- [MINOR] Updated `json_route` and `streaming_json_route` to only parse and serialize requests and responses
|
|
17
|
+
- [MINOR] Added OpenAPI tag ordering, `OpenApiSecurityScheme` declarations, and route security metadata
|
|
18
|
+
- [MINOR] Added OpenAPI security-scheme publishing to authorization decorators
|
|
19
|
+
- [MINOR] Added `authorize_token_request` and `authorize_static_token_request` for token validation
|
|
11
20
|
- [MAJOR] Replace `httpx` with `httpx2` for HTTP requests
|
|
12
21
|
- [MINOR] Use `json_util`/`orjson` for API and database JSON serialization
|
|
13
22
|
- [MINOR] Use Pydantic `JsonValue` for JSON type aliases
|
|
@@ -7,6 +7,8 @@ from pydantic import BaseModel
|
|
|
7
7
|
|
|
8
8
|
from core import logging
|
|
9
9
|
from core.api.api_request import KibaApiRequest
|
|
10
|
+
from core.api.route_metadata import OpenApiSecurityScheme
|
|
11
|
+
from core.api.route_metadata import update_route_metadata
|
|
10
12
|
from core.exceptions import ForbiddenException
|
|
11
13
|
from core.exceptions import UnauthorizedException
|
|
12
14
|
from core.http.basic_authentication import BasicAuthentication
|
|
@@ -41,8 +43,13 @@ async def _authorize_bearer_jwt[ApiRequest: BaseModel](request: KibaApiRequest[A
|
|
|
41
43
|
|
|
42
44
|
def authorize_bearer_jwt[ApiRequest: BaseModel]( # type: ignore[explicit-any]
|
|
43
45
|
authorizer: Authorizer,
|
|
46
|
+
*,
|
|
47
|
+
securityScheme: OpenApiSecurityScheme | None = None,
|
|
44
48
|
) -> typing.Callable[[typing.Callable[[KibaApiRequest[ApiRequest]], _AnyReturn]], typing.Callable[[KibaApiRequest[ApiRequest]], typing.Any]]:
|
|
45
49
|
def decorator(func: typing.Callable[[KibaApiRequest[ApiRequest]], _AnyReturn]) -> typing.Callable[[KibaApiRequest[ApiRequest]], typing.Any]: # type: ignore[explicit-any]
|
|
50
|
+
if securityScheme is not None:
|
|
51
|
+
update_route_metadata(func, {'security': [{securityScheme.name: []}]})
|
|
52
|
+
|
|
46
53
|
@functools.wraps(func)
|
|
47
54
|
async def async_wrapper(request: KibaApiRequest[ApiRequest]) -> typing.Any: # type: ignore[explicit-any, misc]
|
|
48
55
|
request.authJwt = await _authorize_bearer_jwt(request=request, authorizer=authorizer)
|
|
@@ -76,8 +83,13 @@ async def get_basic_authentication_from_authorization_signature[ApiRequest: Base
|
|
|
76
83
|
|
|
77
84
|
def authorize_signature[ApiRequest: BaseModel]( # type: ignore[explicit-any]
|
|
78
85
|
authorizer: SignatureAuthorizer,
|
|
86
|
+
*,
|
|
87
|
+
securityScheme: OpenApiSecurityScheme | None = None,
|
|
79
88
|
) -> typing.Callable[[typing.Callable[[KibaApiRequest[ApiRequest]], _AnyReturn]], typing.Callable[[KibaApiRequest[ApiRequest]], typing.Any]]:
|
|
80
89
|
def decorator(func: typing.Callable[[KibaApiRequest[ApiRequest]], _AnyReturn]) -> typing.Callable[[KibaApiRequest[ApiRequest]], typing.Any]: # type: ignore[explicit-any]
|
|
90
|
+
if securityScheme is not None:
|
|
91
|
+
update_route_metadata(func, {'security': [{securityScheme.name: []}]})
|
|
92
|
+
|
|
81
93
|
@functools.wraps(func)
|
|
82
94
|
async def async_wrapper(request: KibaApiRequest[ApiRequest]) -> typing.Any: # type: ignore[explicit-any, misc]
|
|
83
95
|
request.authBasic = await get_basic_authentication_from_authorization_signature(request=request, authorizer=authorizer)
|
|
@@ -106,18 +118,48 @@ class StaticTokenAuthorizer(TokenAuthorizer):
|
|
|
106
118
|
raise ForbiddenException(message='AUTH_INVALID')
|
|
107
119
|
|
|
108
120
|
|
|
121
|
+
async def authorize_token_request[ApiRequest: BaseModel](request: KibaApiRequest[ApiRequest], authorizer: TokenAuthorizer) -> None:
|
|
122
|
+
authorization = request.headers.get('Authorization')
|
|
123
|
+
if not authorization:
|
|
124
|
+
raise ForbiddenException(message='AUTH_NOT_PROVIDED')
|
|
125
|
+
if not authorization.startswith('Token '):
|
|
126
|
+
raise ForbiddenException(message='AUTH_INVALID')
|
|
127
|
+
await authorizer.validate_token(authorization[6:])
|
|
128
|
+
|
|
129
|
+
|
|
130
|
+
async def authorize_static_token_request[ApiRequest: BaseModel](request: KibaApiRequest[ApiRequest], token: str) -> None:
|
|
131
|
+
authorization = request.headers.get('Authorization')
|
|
132
|
+
if not authorization:
|
|
133
|
+
raise ForbiddenException(message='AUTH_NOT_PROVIDED')
|
|
134
|
+
if not authorization.startswith('Token ') or not hmac.compare_digest(authorization[6:], token):
|
|
135
|
+
raise ForbiddenException(message='AUTH_INVALID')
|
|
136
|
+
|
|
137
|
+
|
|
138
|
+
async def authorize_static_header_request[ApiRequest: BaseModel](
|
|
139
|
+
request: KibaApiRequest[ApiRequest],
|
|
140
|
+
*,
|
|
141
|
+
headerName: str,
|
|
142
|
+
token: str,
|
|
143
|
+
) -> None:
|
|
144
|
+
providedToken = request.headers.get(headerName)
|
|
145
|
+
if providedToken is None:
|
|
146
|
+
raise ForbiddenException(message='AUTH_NOT_PROVIDED')
|
|
147
|
+
if not hmac.compare_digest(providedToken, token):
|
|
148
|
+
raise ForbiddenException(message='AUTH_INVALID')
|
|
149
|
+
|
|
150
|
+
|
|
109
151
|
def authorize_token[ApiRequest: BaseModel]( # type: ignore[explicit-any]
|
|
110
152
|
authorizer: TokenAuthorizer,
|
|
153
|
+
*,
|
|
154
|
+
securityScheme: OpenApiSecurityScheme | None = None,
|
|
111
155
|
) -> typing.Callable[[typing.Callable[[KibaApiRequest[ApiRequest]], _AnyReturn]], typing.Callable[[KibaApiRequest[ApiRequest]], typing.Any]]:
|
|
112
156
|
def decorator(func: typing.Callable[[KibaApiRequest[ApiRequest]], _AnyReturn]) -> typing.Callable[[KibaApiRequest[ApiRequest]], typing.Any]: # type: ignore[explicit-any]
|
|
157
|
+
if securityScheme is not None:
|
|
158
|
+
update_route_metadata(func, {'security': [{securityScheme.name: []}]})
|
|
159
|
+
|
|
113
160
|
@functools.wraps(func)
|
|
114
161
|
async def async_wrapper(request: KibaApiRequest[ApiRequest]) -> typing.Any: # type: ignore[explicit-any, misc]
|
|
115
|
-
|
|
116
|
-
if not authorization:
|
|
117
|
-
raise ForbiddenException(message='AUTH_NOT_PROVIDED')
|
|
118
|
-
if not authorization.startswith('Token '):
|
|
119
|
-
raise ForbiddenException(message='AUTH_INVALID')
|
|
120
|
-
await authorizer.validate_token(authorization[6:])
|
|
162
|
+
await authorize_token_request(request=request, authorizer=authorizer)
|
|
121
163
|
result = func(request)
|
|
122
164
|
# NOTE(krishan711): this is here to support streaming responses which return an async generator
|
|
123
165
|
if hasattr(result, '__aiter__'):
|
|
@@ -37,6 +37,7 @@ def json_route[ApiRequest: BaseModel, ApiResponse: BaseModel](
|
|
|
37
37
|
validationErrorMessage = ', '.join([f'{".".join([str(value) for value in error["loc"]])}: {error["msg"]}' for error in exception.errors()])
|
|
38
38
|
raise BadRequestException(f'Invalid request: {validationErrorMessage}')
|
|
39
39
|
kibaRequest: KibaApiRequest[ApiRequest] = KibaApiRequest(scope=receivedRequest.scope, receive=receivedRequest._receive, send=receivedRequest._send) # noqa: SLF001
|
|
40
|
+
kibaRequest.originIp = typing.cast('str | None', receivedRequest.scope.get('originIp'))
|
|
40
41
|
kibaRequest.data = requestParams
|
|
41
42
|
receivedResponse = await func(kibaRequest)
|
|
42
43
|
if not isinstance(receivedResponse, responseType):
|
|
@@ -0,0 +1,36 @@
|
|
|
1
|
+
from collections.abc import Sequence
|
|
2
|
+
|
|
3
|
+
from starlette.datastructures import Headers
|
|
4
|
+
from starlette.types import ASGIApp
|
|
5
|
+
from starlette.types import Receive
|
|
6
|
+
from starlette.types import Scope
|
|
7
|
+
from starlette.types import Send
|
|
8
|
+
|
|
9
|
+
from core import logging
|
|
10
|
+
|
|
11
|
+
|
|
12
|
+
def _extract_origin_ip(scope: Scope, trustedProxyHeaders: Sequence[str]) -> str | None:
|
|
13
|
+
headers = Headers(scope=scope)
|
|
14
|
+
for headerName in trustedProxyHeaders:
|
|
15
|
+
headerValue = headers.get(headerName)
|
|
16
|
+
if headerValue:
|
|
17
|
+
return headerValue.split(',')[0].strip()
|
|
18
|
+
client = scope.get('client')
|
|
19
|
+
if client:
|
|
20
|
+
return str(client[0])
|
|
21
|
+
return None
|
|
22
|
+
|
|
23
|
+
|
|
24
|
+
class OriginIpMiddleware:
|
|
25
|
+
def __init__(self, app: ASGIApp, trustedProxyHeaders: Sequence[str] = ()) -> None:
|
|
26
|
+
self.app = app
|
|
27
|
+
self.trustedProxyHeaders = trustedProxyHeaders
|
|
28
|
+
if trustedProxyHeaders:
|
|
29
|
+
logging.info(f'OriginIpMiddleware trusting proxy headers for client IP: {", ".join(trustedProxyHeaders)}. This may be spoofable if requests can reach this service directly.')
|
|
30
|
+
|
|
31
|
+
async def __call__(self, scope: Scope, receive: Receive, send: Send) -> None:
|
|
32
|
+
if scope['type'] != 'http':
|
|
33
|
+
await self.app(scope, receive, send)
|
|
34
|
+
return
|
|
35
|
+
scope['originIp'] = _extract_origin_ip(scope=scope, trustedProxyHeaders=self.trustedProxyHeaders)
|
|
36
|
+
await self.app(scope, receive, send)
|
|
@@ -0,0 +1,215 @@
|
|
|
1
|
+
from __future__ import annotations
|
|
2
|
+
|
|
3
|
+
import re
|
|
4
|
+
from collections.abc import Callable
|
|
5
|
+
from collections.abc import Mapping
|
|
6
|
+
from collections.abc import Sequence
|
|
7
|
+
from typing import Protocol
|
|
8
|
+
from typing import cast
|
|
9
|
+
|
|
10
|
+
from pydantic import BaseModel
|
|
11
|
+
from starlette.requests import Request
|
|
12
|
+
from starlette.routing import BaseRoute
|
|
13
|
+
from starlette.schemas import EndpointInfo
|
|
14
|
+
from starlette.schemas import SchemaGenerator
|
|
15
|
+
|
|
16
|
+
from core.api.api_response import KibaJSONResponse
|
|
17
|
+
from core.api.route_metadata import OpenApiSecurityScheme
|
|
18
|
+
from core.api.route_metadata import RouteMetadata
|
|
19
|
+
from core.api.route_metadata import get_route_metadata
|
|
20
|
+
|
|
21
|
+
_PATH_PARAMETER_PATTERN = re.compile(r'{([A-Za-z0-9_]+)(?::[^}]+)?}')
|
|
22
|
+
_RATE_LIMIT_WINDOW_LABELS: dict[str, str] = {
|
|
23
|
+
'perMinute': 'minute',
|
|
24
|
+
'perFiveMinutes': '5 minutes',
|
|
25
|
+
'perHour': 'hour',
|
|
26
|
+
'perDay': 'day',
|
|
27
|
+
}
|
|
28
|
+
|
|
29
|
+
|
|
30
|
+
class OpenApiTag(BaseModel):
|
|
31
|
+
name: str
|
|
32
|
+
description: str | None = None
|
|
33
|
+
|
|
34
|
+
|
|
35
|
+
class OpenApiExtension(Protocol):
|
|
36
|
+
def update_document(self, document: dict[str, object]) -> None: ...
|
|
37
|
+
|
|
38
|
+
def update_operation(self, *, route: EndpointInfo, metadata: RouteMetadata, operation: dict[str, object]) -> None: ...
|
|
39
|
+
|
|
40
|
+
|
|
41
|
+
def _format_rate_limit_text(rateLimit: Mapping[str, int]) -> str:
|
|
42
|
+
parts = [f'{rateLimit[key]} requests/{label}' for key, label in _RATE_LIMIT_WINDOW_LABELS.items() if key in rateLimit]
|
|
43
|
+
return f'Rate limited to {", ".join(parts)}.'
|
|
44
|
+
|
|
45
|
+
|
|
46
|
+
def _schema_reference(model: type[BaseModel], schemas: dict[str, object]) -> dict[str, str]:
|
|
47
|
+
modelSchema = cast(dict[str, object], model.model_json_schema(by_alias=True, ref_template='#/components/schemas/{model}'))
|
|
48
|
+
definitions = modelSchema.pop('$defs', {})
|
|
49
|
+
schemas.update(cast(dict[str, object], definitions))
|
|
50
|
+
schemas[model.__name__] = modelSchema
|
|
51
|
+
return {'$ref': f'#/components/schemas/{model.__name__}'}
|
|
52
|
+
|
|
53
|
+
|
|
54
|
+
def _get_path_parameter_names(path: str) -> set[str]:
|
|
55
|
+
return {match.group(1) for match in _PATH_PARAMETER_PATTERN.finditer(path)}
|
|
56
|
+
|
|
57
|
+
|
|
58
|
+
def _get_request_schema(model: type[BaseModel]) -> dict[str, object]:
|
|
59
|
+
return cast(dict[str, object], model.model_json_schema(by_alias=True))
|
|
60
|
+
|
|
61
|
+
|
|
62
|
+
def _get_path_parameters(model: type[BaseModel], path: str) -> list[dict[str, object]]:
|
|
63
|
+
modelSchema = _get_request_schema(model=model)
|
|
64
|
+
properties = cast(dict[str, dict[str, object]], modelSchema.get('properties', {}))
|
|
65
|
+
pathParameterNames = _get_path_parameter_names(path)
|
|
66
|
+
return [
|
|
67
|
+
{
|
|
68
|
+
'name': name,
|
|
69
|
+
'in': 'path',
|
|
70
|
+
'required': True,
|
|
71
|
+
'schema': properties[name],
|
|
72
|
+
}
|
|
73
|
+
for name in properties
|
|
74
|
+
if name in pathParameterNames
|
|
75
|
+
]
|
|
76
|
+
|
|
77
|
+
|
|
78
|
+
def _get_query_parameters(model: type[BaseModel], path: str) -> list[dict[str, object]]:
|
|
79
|
+
modelSchema = _get_request_schema(model=model)
|
|
80
|
+
properties = cast(dict[str, dict[str, object]], modelSchema.get('properties', {}))
|
|
81
|
+
required = set(cast(list[str], modelSchema.get('required', [])))
|
|
82
|
+
pathParameterNames = _get_path_parameter_names(path)
|
|
83
|
+
return [
|
|
84
|
+
{
|
|
85
|
+
'name': name,
|
|
86
|
+
'in': 'query',
|
|
87
|
+
'required': name in required,
|
|
88
|
+
'schema': propertySchema,
|
|
89
|
+
}
|
|
90
|
+
for name, propertySchema in properties.items()
|
|
91
|
+
if name not in pathParameterNames
|
|
92
|
+
]
|
|
93
|
+
|
|
94
|
+
|
|
95
|
+
def _get_request_body_schema(model: type[BaseModel], path: str, schemas: dict[str, object]) -> dict[str, object]:
|
|
96
|
+
modelSchema = cast(dict[str, object], model.model_json_schema(by_alias=True, ref_template='#/components/schemas/{model}'))
|
|
97
|
+
definitions = modelSchema.pop('$defs', {})
|
|
98
|
+
schemas.update(cast(dict[str, object], definitions))
|
|
99
|
+
properties = cast(dict[str, dict[str, object]], modelSchema.get('properties', {}))
|
|
100
|
+
pathParameterNames = _get_path_parameter_names(path)
|
|
101
|
+
modelSchema['properties'] = {name: schema for name, schema in properties.items() if name not in pathParameterNames}
|
|
102
|
+
required = cast(list[str], modelSchema.get('required', []))
|
|
103
|
+
modelSchema['required'] = [name for name in required if name not in pathParameterNames]
|
|
104
|
+
return modelSchema
|
|
105
|
+
|
|
106
|
+
|
|
107
|
+
class OpenApiSchemaGenerator(SchemaGenerator):
|
|
108
|
+
def __init__(
|
|
109
|
+
self,
|
|
110
|
+
title: str,
|
|
111
|
+
version: str,
|
|
112
|
+
description: str,
|
|
113
|
+
tags: Sequence[OpenApiTag],
|
|
114
|
+
securitySchemes: Sequence[OpenApiSecurityScheme],
|
|
115
|
+
extensions: Sequence[OpenApiExtension] = (),
|
|
116
|
+
) -> None:
|
|
117
|
+
super().__init__({'openapi': '3.0.3', 'info': {'title': title, 'version': version}})
|
|
118
|
+
self.title = title
|
|
119
|
+
self.version = version
|
|
120
|
+
self.description = description
|
|
121
|
+
self.tags = tags
|
|
122
|
+
self.securitySchemes = {scheme.name: scheme.definition for scheme in securitySchemes}
|
|
123
|
+
self.tagOrder = {tag.name: index for index, tag in enumerate(tags)}
|
|
124
|
+
self.extensions = extensions
|
|
125
|
+
|
|
126
|
+
def get_schema(self, routes: list[BaseRoute]) -> dict[str, object]:
|
|
127
|
+
document: dict[str, object] = {
|
|
128
|
+
'openapi': '3.0.3',
|
|
129
|
+
'info': {
|
|
130
|
+
'title': self.title,
|
|
131
|
+
'version': self.version,
|
|
132
|
+
'description': self.description,
|
|
133
|
+
},
|
|
134
|
+
'tags': [tag.model_dump(exclude_none=True) for tag in self.tags],
|
|
135
|
+
'paths': {},
|
|
136
|
+
'components': {
|
|
137
|
+
'schemas': {},
|
|
138
|
+
'securitySchemes': self.securitySchemes,
|
|
139
|
+
},
|
|
140
|
+
}
|
|
141
|
+
paths = cast(dict[str, object], document['paths'])
|
|
142
|
+
components = cast(dict[str, object], document['components'])
|
|
143
|
+
schemas = cast(dict[str, object], components['schemas'])
|
|
144
|
+
|
|
145
|
+
def sort_key(endpoint: object) -> int:
|
|
146
|
+
metadata = get_route_metadata(getattr(endpoint, 'func', None))
|
|
147
|
+
if metadata.get('operationId') is None:
|
|
148
|
+
return len(self.tagOrder)
|
|
149
|
+
endpointTags = metadata.get('tags', [])
|
|
150
|
+
tag = endpointTags[0] if endpointTags else ''
|
|
151
|
+
return self.tagOrder.get(tag, len(self.tagOrder))
|
|
152
|
+
|
|
153
|
+
for endpoint in sorted(self.get_endpoints(routes), key=sort_key):
|
|
154
|
+
metadata = get_route_metadata(endpoint.func)
|
|
155
|
+
if metadata.get('operationId') is None:
|
|
156
|
+
continue
|
|
157
|
+
requestType = metadata['requestType']
|
|
158
|
+
responseType = metadata['responseType']
|
|
159
|
+
operation: dict[str, object] = {
|
|
160
|
+
'operationId': metadata['operationId'],
|
|
161
|
+
'summary': metadata['summary'],
|
|
162
|
+
'tags': metadata['tags'],
|
|
163
|
+
'responses': {
|
|
164
|
+
'200': {
|
|
165
|
+
'description': 'Successful response',
|
|
166
|
+
'content': {
|
|
167
|
+
'application/x-ndjson' if metadata['streamed'] else 'application/json': {
|
|
168
|
+
'schema': _schema_reference(model=responseType, schemas=schemas),
|
|
169
|
+
},
|
|
170
|
+
},
|
|
171
|
+
},
|
|
172
|
+
},
|
|
173
|
+
}
|
|
174
|
+
if metadata['description'] is not None:
|
|
175
|
+
operation['description'] = metadata['description']
|
|
176
|
+
rateLimit = metadata.get('rateLimit')
|
|
177
|
+
if rateLimit:
|
|
178
|
+
rateLimitText = _format_rate_limit_text(rateLimit=cast(Mapping[str, int], rateLimit))
|
|
179
|
+
existingDescription = cast(str, operation.get('description', ''))
|
|
180
|
+
operation['description'] = f'{existingDescription}\n\n{rateLimitText}'.strip()
|
|
181
|
+
security = metadata.get('security')
|
|
182
|
+
if security is not None:
|
|
183
|
+
operation['security'] = security
|
|
184
|
+
if endpoint.http_method in {'get', 'delete'}:
|
|
185
|
+
parameters = _get_path_parameters(model=requestType, path=endpoint.path) + _get_query_parameters(model=requestType, path=endpoint.path)
|
|
186
|
+
if parameters:
|
|
187
|
+
operation['parameters'] = parameters
|
|
188
|
+
else:
|
|
189
|
+
pathParameters = _get_path_parameters(model=requestType, path=endpoint.path)
|
|
190
|
+
if pathParameters:
|
|
191
|
+
operation['parameters'] = pathParameters
|
|
192
|
+
bodySchema = _get_request_body_schema(model=requestType, path=endpoint.path, schemas=schemas)
|
|
193
|
+
if bodySchema.get('properties'):
|
|
194
|
+
operation['requestBody'] = {
|
|
195
|
+
'required': bool(bodySchema.get('required')),
|
|
196
|
+
'content': {
|
|
197
|
+
'application/json': {
|
|
198
|
+
'schema': bodySchema,
|
|
199
|
+
},
|
|
200
|
+
},
|
|
201
|
+
}
|
|
202
|
+
for extension in self.extensions:
|
|
203
|
+
extension.update_operation(route=endpoint, metadata=metadata, operation=operation)
|
|
204
|
+
pathOperations = cast(dict[str, object], paths.setdefault(endpoint.path, {}))
|
|
205
|
+
pathOperations[endpoint.http_method] = operation
|
|
206
|
+
for extension in self.extensions:
|
|
207
|
+
extension.update_document(document=document)
|
|
208
|
+
return document
|
|
209
|
+
|
|
210
|
+
|
|
211
|
+
def create_openapi_response(schemaGenerator: OpenApiSchemaGenerator) -> Callable[[Request], KibaJSONResponse]:
|
|
212
|
+
def get_openapi_response(request: Request) -> KibaJSONResponse:
|
|
213
|
+
return KibaJSONResponse(content=schemaGenerator.get_schema(routes=request.app.routes))
|
|
214
|
+
|
|
215
|
+
return get_openapi_response
|
|
@@ -0,0 +1,110 @@
|
|
|
1
|
+
import functools
|
|
2
|
+
import time
|
|
3
|
+
import typing
|
|
4
|
+
from collections.abc import AsyncIterator
|
|
5
|
+
|
|
6
|
+
from core.api.api_request import KibaApiRequest
|
|
7
|
+
from core.api.route_metadata import RateLimitConfig
|
|
8
|
+
from core.api.route_metadata import update_route_metadata
|
|
9
|
+
from core.exceptions import InternalServerErrorException
|
|
10
|
+
from core.exceptions import TooManyRequestsException
|
|
11
|
+
|
|
12
|
+
_AnyReturn = typing.Awaitable[typing.Any] | AsyncIterator[typing.Any] # type: ignore[explicit-any]
|
|
13
|
+
_SWEEP_INTERVAL = 1000
|
|
14
|
+
_MAX_ENTRY_AGE_SECONDS = 25 * 60 * 60
|
|
15
|
+
|
|
16
|
+
|
|
17
|
+
class _WindowConfig(typing.NamedTuple):
|
|
18
|
+
label: str
|
|
19
|
+
limit: int
|
|
20
|
+
windowSeconds: int
|
|
21
|
+
|
|
22
|
+
|
|
23
|
+
class _WindowState:
|
|
24
|
+
__slots__ = ('count', 'windowStart')
|
|
25
|
+
|
|
26
|
+
def __init__(self, windowStart: float) -> None:
|
|
27
|
+
self.count = 0
|
|
28
|
+
self.windowStart = windowStart
|
|
29
|
+
|
|
30
|
+
|
|
31
|
+
_store: dict[str, dict[str, _WindowState]] = {}
|
|
32
|
+
_checkCount = 0
|
|
33
|
+
|
|
34
|
+
|
|
35
|
+
def _resolve_user_identity(request: KibaApiRequest[typing.Any]) -> str: # type: ignore[explicit-any]
|
|
36
|
+
if request.authBasic is not None:
|
|
37
|
+
return request.authBasic.username
|
|
38
|
+
if request.authJwt is not None:
|
|
39
|
+
subject = request.authJwt.payloadDict.get('sub')
|
|
40
|
+
if subject:
|
|
41
|
+
return str(subject)
|
|
42
|
+
raise InternalServerErrorException(message='rate_limit(keyBy="user") requires an auth decorator to run first')
|
|
43
|
+
|
|
44
|
+
|
|
45
|
+
def _resolve_ip_identity(request: KibaApiRequest[typing.Any]) -> str: # type: ignore[explicit-any]
|
|
46
|
+
if request.originIp is None:
|
|
47
|
+
raise InternalServerErrorException(message='rate_limit(keyBy="ip") requires OriginIpMiddleware to be installed')
|
|
48
|
+
return request.originIp
|
|
49
|
+
|
|
50
|
+
|
|
51
|
+
def _sweep_expired_entries(now: float) -> None:
|
|
52
|
+
staleKeys = [storeKey for storeKey, windowStates in _store.items() if all((now - windowState.windowStart) >= _MAX_ENTRY_AGE_SECONDS for windowState in windowStates.values())]
|
|
53
|
+
for storeKey in staleKeys:
|
|
54
|
+
del _store[storeKey]
|
|
55
|
+
|
|
56
|
+
|
|
57
|
+
def check_rate_limit(*, request: KibaApiRequest[typing.Any], routeKey: str, rateLimit: RateLimitConfig) -> None: # type: ignore[explicit-any]
|
|
58
|
+
windows: list[_WindowConfig] = []
|
|
59
|
+
if 'perMinute' in rateLimit:
|
|
60
|
+
windows.append(_WindowConfig(label='perMinute', limit=rateLimit['perMinute'], windowSeconds=60))
|
|
61
|
+
if 'perFiveMinutes' in rateLimit:
|
|
62
|
+
windows.append(_WindowConfig(label='perFiveMinutes', limit=rateLimit['perFiveMinutes'], windowSeconds=5 * 60))
|
|
63
|
+
if 'perHour' in rateLimit:
|
|
64
|
+
windows.append(_WindowConfig(label='perHour', limit=rateLimit['perHour'], windowSeconds=60 * 60))
|
|
65
|
+
if 'perDay' in rateLimit:
|
|
66
|
+
windows.append(_WindowConfig(label='perDay', limit=rateLimit['perDay'], windowSeconds=24 * 60 * 60))
|
|
67
|
+
if not windows:
|
|
68
|
+
raise ValueError('rate limit requires at least one of perMinute, perFiveMinutes, perHour, perDay')
|
|
69
|
+
global _checkCount # noqa: PLW0603
|
|
70
|
+
keyBy = rateLimit.get('keyBy', 'user')
|
|
71
|
+
identity = _resolve_user_identity(request) if keyBy == 'user' else _resolve_ip_identity(request)
|
|
72
|
+
storeKey = f'{routeKey}:{identity}'
|
|
73
|
+
now = time.monotonic()
|
|
74
|
+
_checkCount += 1
|
|
75
|
+
if _checkCount % _SWEEP_INTERVAL == 0:
|
|
76
|
+
_sweep_expired_entries(now=now)
|
|
77
|
+
windowStates = _store.setdefault(storeKey, {})
|
|
78
|
+
retryAfterSeconds = 0
|
|
79
|
+
for window in windows:
|
|
80
|
+
windowState = windowStates.get(window.label)
|
|
81
|
+
if windowState is None or (now - windowState.windowStart) >= window.windowSeconds:
|
|
82
|
+
windowState = _WindowState(windowStart=now)
|
|
83
|
+
windowStates[window.label] = windowState
|
|
84
|
+
if windowState.count >= window.limit:
|
|
85
|
+
remainingSeconds = window.windowSeconds - (now - windowState.windowStart)
|
|
86
|
+
retryAfterSeconds = max(retryAfterSeconds, int(remainingSeconds) + 1)
|
|
87
|
+
if retryAfterSeconds > 0:
|
|
88
|
+
raise TooManyRequestsException(message='RATE_LIMITED', retryAfterSeconds=retryAfterSeconds)
|
|
89
|
+
for window in windows:
|
|
90
|
+
windowStates[window.label].count += 1
|
|
91
|
+
|
|
92
|
+
|
|
93
|
+
def rate_limit( # type: ignore[explicit-any]
|
|
94
|
+
rateLimit: RateLimitConfig,
|
|
95
|
+
) -> typing.Callable[[typing.Callable[[KibaApiRequest[typing.Any]], _AnyReturn]], typing.Callable[[KibaApiRequest[typing.Any]], typing.Any]]:
|
|
96
|
+
def decorator(func: typing.Callable[[KibaApiRequest[typing.Any]], _AnyReturn]) -> typing.Callable[[KibaApiRequest[typing.Any]], typing.Any]: # type: ignore[explicit-any]
|
|
97
|
+
update_route_metadata(func, {'rateLimit': rateLimit})
|
|
98
|
+
routeKey = getattr(func, '__qualname__', type(func).__name__)
|
|
99
|
+
|
|
100
|
+
@functools.wraps(func)
|
|
101
|
+
async def async_wrapper(request: KibaApiRequest[typing.Any]) -> typing.Any: # type: ignore[explicit-any, misc]
|
|
102
|
+
check_rate_limit(request=request, routeKey=routeKey, rateLimit=rateLimit)
|
|
103
|
+
result = func(request)
|
|
104
|
+
if hasattr(result, '__aiter__'):
|
|
105
|
+
return result
|
|
106
|
+
return await result
|
|
107
|
+
|
|
108
|
+
return async_wrapper
|
|
109
|
+
|
|
110
|
+
return decorator
|
|
@@ -0,0 +1,52 @@
|
|
|
1
|
+
import contextvars
|
|
2
|
+
import typing
|
|
3
|
+
from collections.abc import Callable
|
|
4
|
+
from contextlib import contextmanager
|
|
5
|
+
from dataclasses import dataclass
|
|
6
|
+
|
|
7
|
+
from starlette.types import ASGIApp
|
|
8
|
+
from starlette.types import Receive
|
|
9
|
+
from starlette.types import Scope
|
|
10
|
+
from starlette.types import Send
|
|
11
|
+
|
|
12
|
+
|
|
13
|
+
@dataclass
|
|
14
|
+
class RequestContext:
|
|
15
|
+
originIp: str | None
|
|
16
|
+
|
|
17
|
+
|
|
18
|
+
def create_request_context(scope: Scope) -> RequestContext:
|
|
19
|
+
return RequestContext(originIp=typing.cast('str | None', scope.get('originIp')))
|
|
20
|
+
|
|
21
|
+
|
|
22
|
+
class RequestContextHolder[RequestContextType: RequestContext]:
|
|
23
|
+
def __init__(self) -> None:
|
|
24
|
+
self._valueContext = contextvars.ContextVar[RequestContextType | None]('_valueContext', default=None)
|
|
25
|
+
|
|
26
|
+
def get_value(self) -> RequestContextType:
|
|
27
|
+
value = self._valueContext.get()
|
|
28
|
+
if value is None:
|
|
29
|
+
raise RuntimeError('No request context is active')
|
|
30
|
+
return value
|
|
31
|
+
|
|
32
|
+
@contextmanager
|
|
33
|
+
def use_value(self, value: RequestContextType) -> typing.Iterator[RequestContextType]:
|
|
34
|
+
token = self._valueContext.set(value)
|
|
35
|
+
try:
|
|
36
|
+
yield value
|
|
37
|
+
finally:
|
|
38
|
+
self._valueContext.reset(token)
|
|
39
|
+
|
|
40
|
+
|
|
41
|
+
class RequestContextMiddleware:
|
|
42
|
+
def __init__(self, app: ASGIApp, requestContextHolder: RequestContextHolder[typing.Any], requestContextFactory: Callable[[Scope], RequestContext]) -> None: # type: ignore[explicit-any]
|
|
43
|
+
self.app = app
|
|
44
|
+
self.requestContextHolder = requestContextHolder
|
|
45
|
+
self.requestContextFactory = requestContextFactory
|
|
46
|
+
|
|
47
|
+
async def __call__(self, scope: Scope, receive: Receive, send: Send) -> None:
|
|
48
|
+
if scope['type'] != 'http':
|
|
49
|
+
await self.app(scope, receive, send)
|
|
50
|
+
return
|
|
51
|
+
with self.requestContextHolder.use_value(self.requestContextFactory(scope)):
|
|
52
|
+
await self.app(scope, receive, send)
|
|
@@ -0,0 +1,80 @@
|
|
|
1
|
+
import functools
|
|
2
|
+
import typing
|
|
3
|
+
from collections.abc import AsyncIterator
|
|
4
|
+
from collections.abc import Mapping
|
|
5
|
+
|
|
6
|
+
from pydantic import BaseModel
|
|
7
|
+
from starlette.requests import Request
|
|
8
|
+
from starlette.responses import StreamingResponse
|
|
9
|
+
|
|
10
|
+
from core.api.api_request import KibaApiRequest
|
|
11
|
+
from core.api.api_response import KibaJSONResponse
|
|
12
|
+
from core.api.json_route import json_route
|
|
13
|
+
from core.api.rate_limit import rate_limit
|
|
14
|
+
from core.api.route_auth import RouteAuthResolver
|
|
15
|
+
from core.api.route_metadata import RateLimitConfig
|
|
16
|
+
from core.api.route_metadata import RouteMetadata
|
|
17
|
+
from core.api.route_metadata import update_route_metadata
|
|
18
|
+
from core.api.streaming_json_route import streaming_json_route
|
|
19
|
+
|
|
20
|
+
_AnyReturn = typing.Awaitable[typing.Any] | AsyncIterator[typing.Any] # type: ignore[explicit-any]
|
|
21
|
+
|
|
22
|
+
|
|
23
|
+
def _authorize_route[ApiRequest: BaseModel](
|
|
24
|
+
authResolver: RouteAuthResolver,
|
|
25
|
+
auth: str,
|
|
26
|
+
) -> typing.Callable[[typing.Callable[[KibaApiRequest[ApiRequest]], _AnyReturn]], typing.Callable[[KibaApiRequest[ApiRequest]], typing.Any]]:
|
|
27
|
+
def decorator(func: typing.Callable[[KibaApiRequest[ApiRequest]], _AnyReturn]) -> typing.Callable[[KibaApiRequest[ApiRequest]], typing.Any]: # type: ignore[explicit-any]
|
|
28
|
+
@functools.wraps(func)
|
|
29
|
+
async def async_wrapper(request: KibaApiRequest[ApiRequest]) -> typing.Any: # type: ignore[explicit-any, misc]
|
|
30
|
+
await authResolver.authorize_route(auth=auth, request=request)
|
|
31
|
+
result = func(request)
|
|
32
|
+
if hasattr(result, '__aiter__'):
|
|
33
|
+
return result
|
|
34
|
+
return await result
|
|
35
|
+
|
|
36
|
+
return async_wrapper
|
|
37
|
+
|
|
38
|
+
return decorator
|
|
39
|
+
|
|
40
|
+
|
|
41
|
+
def route[ApiRequest: BaseModel, ApiResponse: BaseModel]( # type: ignore[explicit-any]
|
|
42
|
+
requestType: typing.Type[ApiRequest],
|
|
43
|
+
responseType: typing.Type[ApiResponse],
|
|
44
|
+
*,
|
|
45
|
+
authResolver: RouteAuthResolver,
|
|
46
|
+
isStreaming: bool = False,
|
|
47
|
+
operationId: str | None = None,
|
|
48
|
+
summary: str | None = None,
|
|
49
|
+
description: str | None = None,
|
|
50
|
+
tags: list[str] | None = None,
|
|
51
|
+
auth: str | None = None,
|
|
52
|
+
rateLimit: RateLimitConfig | None = None,
|
|
53
|
+
extensions: Mapping[str, object] | None = None,
|
|
54
|
+
) -> typing.Callable[[typing.Callable[[KibaApiRequest[ApiRequest]], _AnyReturn]], typing.Callable[[Request], typing.Awaitable[KibaJSONResponse | StreamingResponse]]]:
|
|
55
|
+
securitySchemeNames = authResolver.get_route_security_schemes(auth=auth) if auth is not None else []
|
|
56
|
+
|
|
57
|
+
def decorator(func: typing.Callable[[KibaApiRequest[ApiRequest]], _AnyReturn]) -> typing.Callable[[Request], typing.Awaitable[KibaJSONResponse | StreamingResponse]]: # type: ignore[explicit-any]
|
|
58
|
+
handler: typing.Callable[[KibaApiRequest[ApiRequest]], _AnyReturn] = func
|
|
59
|
+
if rateLimit is not None:
|
|
60
|
+
handler = rate_limit(rateLimit)(handler)
|
|
61
|
+
if auth is not None:
|
|
62
|
+
handler = _authorize_route(authResolver, auth)(handler)
|
|
63
|
+
endpointDecorator = streaming_json_route if isStreaming else json_route
|
|
64
|
+
endpoint = endpointDecorator(requestType=requestType, responseType=responseType)(handler) # type: ignore[arg-type, ty:invalid-argument-type]
|
|
65
|
+
metadata: RouteMetadata = {
|
|
66
|
+
'requestType': requestType,
|
|
67
|
+
'responseType': responseType,
|
|
68
|
+
'streamed': isStreaming,
|
|
69
|
+
'operationId': operationId,
|
|
70
|
+
'summary': summary,
|
|
71
|
+
'description': description,
|
|
72
|
+
'tags': tags or [],
|
|
73
|
+
'extensions': dict(extensions) if extensions else {},
|
|
74
|
+
}
|
|
75
|
+
update_route_metadata(endpoint, metadata)
|
|
76
|
+
if securitySchemeNames:
|
|
77
|
+
update_route_metadata(endpoint, {'security': [{name: []} for name in securitySchemeNames]})
|
|
78
|
+
return endpoint
|
|
79
|
+
|
|
80
|
+
return decorator
|
|
@@ -0,0 +1,11 @@
|
|
|
1
|
+
from typing import Protocol
|
|
2
|
+
|
|
3
|
+
from pydantic import BaseModel
|
|
4
|
+
|
|
5
|
+
from core.api.api_request import KibaApiRequest
|
|
6
|
+
|
|
7
|
+
|
|
8
|
+
class RouteAuthResolver(Protocol):
|
|
9
|
+
async def authorize_route[ApiRequest: BaseModel](self, *, auth: str, request: KibaApiRequest[ApiRequest]) -> None: ...
|
|
10
|
+
|
|
11
|
+
def get_route_security_schemes(self, *, auth: str) -> list[str]: ...
|