reqkey 0.1.0__py3-none-any.whl

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.
reqkey/asgi.py ADDED
@@ -0,0 +1,612 @@
1
+ """Generic ASGI integration for ReqKey."""
2
+
3
+ from __future__ import annotations
4
+
5
+ import inspect
6
+ import logging
7
+ import time
8
+ import zlib
9
+ from collections.abc import Awaitable, Callable, Iterable, Mapping
10
+ from datetime import datetime
11
+ from typing import Any, Literal, Protocol
12
+ from urllib.parse import urlencode
13
+
14
+ from starlette.datastructures import Headers, MutableHeaders
15
+ from starlette.requests import Request
16
+ from starlette.responses import JSONResponse
17
+ from starlette.types import ASGIApp, Message, Receive, Scope, Send
18
+
19
+ from .client import DEFAULT_BASE_URL, DEFAULT_TIMEOUT_SECONDS, AsyncReqKey
20
+ from .exceptions import ReqKeyConfigurationError, ReqKeyError
21
+ from .models import VerificationReason, VerificationResult
22
+
23
+ logger = logging.getLogger("reqkey")
24
+
25
+ MAX_RESPONSE_BODY_CHARACTERS = 1000
26
+ _MAX_CAPTURE_BYTES = MAX_RESPONSE_BODY_CHARACTERS * 4 + 4
27
+
28
+ Mode = Literal["validate", "ingest", "both"]
29
+ FailureMode = Literal["closed", "open"]
30
+ KeyLocation = Literal["header", "query", "cookie"]
31
+ KeyScheme = Literal["raw", "bearer"]
32
+ CreditsResolver = Callable[[Request], int | Awaitable[int]]
33
+ ProtectionResolver = Callable[[Request], bool | Awaitable[bool]]
34
+ ConsumerKeyResolver = Callable[[Request], str | None | Awaitable[str | None]]
35
+ RequestIdResolver = Callable[[Request], str | None | Awaitable[str | None]]
36
+ PathResolver = Callable[[Request], str | Awaitable[str]]
37
+
38
+ DEFAULT_ERROR_MESSAGES: dict[str, str] = {
39
+ "missing_api_key": "An API key is required.",
40
+ "invalid_api_key": "The API key is invalid or inactive.",
41
+ "insufficient_credits": "The API key has insufficient credits.",
42
+ "access_denied": "The API key is not allowed to access this API.",
43
+ "rate_limited": "The API key has exceeded its rate limit.",
44
+ "reqkey_unavailable": "API key verification is temporarily unavailable.",
45
+ }
46
+
47
+ DEFAULT_EXCLUDED_HEADERS = frozenset(
48
+ {
49
+ "authorization",
50
+ "cookie",
51
+ "proxy-authorization",
52
+ "set-cookie",
53
+ "x-api-key",
54
+ }
55
+ )
56
+
57
+
58
+ class _ResponseBodyCapture:
59
+ """Keep at most the bytes needed for the first 1,000 decoded characters."""
60
+
61
+ def __init__(self, *, enabled: bool, headers: Headers) -> None:
62
+ content_type = headers.get("content-type", "").lower()
63
+ self.enabled = enabled and (
64
+ not content_type
65
+ or any(
66
+ marker in content_type
67
+ for marker in (
68
+ "json",
69
+ "text/",
70
+ "xml",
71
+ "javascript",
72
+ "x-www-form-urlencoded",
73
+ )
74
+ )
75
+ )
76
+ self._body = bytearray()
77
+ self._decompressor = (
78
+ zlib.decompressobj(16 + zlib.MAX_WBITS)
79
+ if self.enabled and headers.get("content-encoding", "").lower() == "gzip"
80
+ else None
81
+ )
82
+
83
+ def feed(self, chunk: bytes) -> None:
84
+ if not self.enabled or not chunk or len(self._body) >= _MAX_CAPTURE_BYTES:
85
+ return
86
+ remaining = _MAX_CAPTURE_BYTES - len(self._body)
87
+ if self._decompressor is None:
88
+ self._body.extend(chunk[:remaining])
89
+ return
90
+ try:
91
+ self._body.extend(self._decompressor.decompress(chunk, remaining))
92
+ except zlib.error:
93
+ self.enabled = False
94
+ self._body.clear()
95
+ logger.debug("Could not decompress a gzip response for ReqKey analytics.")
96
+
97
+ def text(self) -> str | None:
98
+ if not self.enabled:
99
+ return None
100
+ if self._decompressor is not None and len(self._body) < _MAX_CAPTURE_BYTES:
101
+ try:
102
+ remaining = _MAX_CAPTURE_BYTES - len(self._body)
103
+ self._body.extend(self._decompressor.flush(remaining))
104
+ except zlib.error:
105
+ logger.debug("Could not finish gzip response capture for ReqKey analytics.")
106
+ return self._body.decode("utf-8", errors="replace")[:MAX_RESPONSE_BODY_CHARACTERS]
107
+
108
+
109
+ class AsyncReqKeyLike(Protocol):
110
+ async def verify(
111
+ self,
112
+ key: str,
113
+ *,
114
+ api_id: str | None = None,
115
+ credits: int = 1,
116
+ resource: str | None = None,
117
+ ) -> VerificationResult: ...
118
+
119
+ async def ingest(
120
+ self,
121
+ request_id: str | None = None,
122
+ *,
123
+ api_id: str | None = None,
124
+ method: str | None = None,
125
+ endpoint: str | None = None,
126
+ path: str | None = None,
127
+ status_code: int | None = None,
128
+ latency_ms: int | None = None,
129
+ client_ip: str | None = None,
130
+ user_agent: str | None = None,
131
+ user_id: str | None = None,
132
+ query_params: Mapping[str, Any] | None = None,
133
+ request_headers: Mapping[str, str] | None = None,
134
+ response_headers: Mapping[str, str] | None = None,
135
+ request_body: str | None = None,
136
+ response_body: str | None = None,
137
+ timestamp: datetime | str | None = None,
138
+ ) -> None: ...
139
+
140
+
141
+ class ReqKeyMiddleware:
142
+ """Pure ASGI middleware for ASGI 3 applications.
143
+
144
+ `mode="validate"` gates requests only, `mode="ingest"` records traffic
145
+ only, and `mode="both"` validates before the handler and then performs a
146
+ correlated, awaited ingestion before releasing the response.
147
+ """
148
+
149
+ def __init__(
150
+ self,
151
+ app: ASGIApp,
152
+ *,
153
+ api_id: str,
154
+ project_key: str | None = None,
155
+ root_key: str | None = None,
156
+ client: AsyncReqKeyLike | None = None,
157
+ base_url: str = DEFAULT_BASE_URL,
158
+ timeout: float = DEFAULT_TIMEOUT_SECONDS,
159
+ mode: Mode = "both",
160
+ enabled: bool = True,
161
+ key_location: KeyLocation = "header",
162
+ key_name: str = "X-API-Key",
163
+ key_scheme: KeyScheme = "raw",
164
+ get_consumer_key: ConsumerKeyResolver | None = None,
165
+ credits: int | CreditsResolver = 1,
166
+ exclude_paths: Iterable[str] = (),
167
+ skip_methods: Iterable[str] = ("OPTIONS",),
168
+ should_protect: ProtectionResolver | None = None,
169
+ request_id_resolver: RequestIdResolver | None = None,
170
+ path_resolver: PathResolver | None = None,
171
+ error_messages: Mapping[str, str] | None = None,
172
+ capture_query_params: bool = False,
173
+ capture_request_headers: bool = False,
174
+ capture_response_headers: bool = False,
175
+ capture_response_body: bool = False,
176
+ capture_client_ip: bool = False,
177
+ capture_user_agent: bool = True,
178
+ excluded_headers: Iterable[str] = (),
179
+ failure_mode: FailureMode = "closed",
180
+ ) -> None:
181
+ if not api_id.strip():
182
+ raise ReqKeyConfigurationError("api_id cannot be empty.")
183
+ if client is not None and (project_key is not None or root_key is not None):
184
+ raise ReqKeyConfigurationError(
185
+ "Pass client or project_key/root_key, not both."
186
+ )
187
+ if mode not in {"validate", "ingest", "both"}:
188
+ raise ReqKeyConfigurationError(
189
+ "mode must be 'validate', 'ingest', or 'both'."
190
+ )
191
+ if key_location not in {"header", "query", "cookie"}:
192
+ raise ReqKeyConfigurationError(
193
+ "key_location must be 'header', 'query', or 'cookie'."
194
+ )
195
+ if key_scheme not in {"raw", "bearer"}:
196
+ raise ReqKeyConfigurationError("key_scheme must be 'raw' or 'bearer'.")
197
+ if key_location != "header" and key_scheme != "raw":
198
+ raise ReqKeyConfigurationError(
199
+ "Query-parameter and cookie keys must use the raw scheme."
200
+ )
201
+ if failure_mode not in {"closed", "open"}:
202
+ raise ReqKeyConfigurationError("failure_mode must be 'closed' or 'open'.")
203
+ if not callable(credits) and (
204
+ isinstance(credits, bool) or not isinstance(credits, int) or credits < 0
205
+ ):
206
+ raise ReqKeyConfigurationError(
207
+ "credits must be a non-negative integer or a callable."
208
+ )
209
+ self.app = app
210
+ self.api_id = api_id
211
+ self.client = client or AsyncReqKey(
212
+ project_key=project_key,
213
+ root_key=root_key,
214
+ base_url=base_url,
215
+ timeout=timeout,
216
+ )
217
+ self._owns_client = client is None
218
+ self.mode = mode
219
+ self.enabled = enabled
220
+ self.key_location = key_location
221
+ self.key_name = key_name
222
+ self.key_scheme = key_scheme
223
+ self.get_consumer_key = get_consumer_key
224
+ self.credits = credits
225
+ self.exclude_paths = tuple(exclude_paths)
226
+ self.skip_methods = frozenset(method.upper() for method in skip_methods)
227
+ self.should_protect = should_protect
228
+ self.request_id_resolver = request_id_resolver
229
+ self.path_resolver = path_resolver
230
+ self.error_messages = {**DEFAULT_ERROR_MESSAGES, **dict(error_messages or {})}
231
+ self.capture_query_params = capture_query_params
232
+ self.capture_request_headers = capture_request_headers
233
+ self.capture_response_headers = capture_response_headers
234
+ self.capture_response_body = capture_response_body
235
+ self.capture_client_ip = capture_client_ip
236
+ self.capture_user_agent = capture_user_agent
237
+ self.excluded_headers = DEFAULT_EXCLUDED_HEADERS | {
238
+ header.lower() for header in excluded_headers
239
+ } | {key_name.lower()}
240
+ self.failure_mode = failure_mode
241
+
242
+ async def __call__(self, scope: Scope, receive: Receive, send: Send) -> None:
243
+ if scope["type"] == "lifespan":
244
+ try:
245
+ await self.app(scope, receive, send)
246
+ finally:
247
+ if self._owns_client and isinstance(self.client, AsyncReqKey):
248
+ await self.client.close()
249
+ return
250
+
251
+ if scope["type"] != "http":
252
+ await self.app(scope, receive, send)
253
+ return
254
+
255
+ scope = dict(scope)
256
+ scope["state"] = dict(scope.get("state", {}))
257
+ request = Request(scope, receive=receive)
258
+ if not self.enabled or not await self._applies_to(request):
259
+ await self.app(scope, receive, send)
260
+ return
261
+
262
+ decision: VerificationResult | None = None
263
+ validation_time_ms: float | None = None
264
+
265
+ if self.mode in {"validate", "both"}:
266
+ consumer_key = await self._consumer_key(request)
267
+ if consumer_key is None:
268
+ await self._respond(
269
+ scope,
270
+ receive,
271
+ send,
272
+ status_code=401,
273
+ error="missing_api_key",
274
+ )
275
+ return
276
+
277
+ validation_started = time.perf_counter()
278
+ try:
279
+ decision = await self.client.verify(
280
+ consumer_key,
281
+ api_id=self.api_id,
282
+ credits=await self._credit_cost(request),
283
+ resource=await self._resource_path(request),
284
+ )
285
+ validation_time_ms = (time.perf_counter() - validation_started) * 1000
286
+ except ReqKeyError as exc:
287
+ validation_time_ms = (time.perf_counter() - validation_started) * 1000
288
+ if self.failure_mode == "open":
289
+ scope.setdefault("state", {})["reqkey_error"] = exc
290
+ else:
291
+ logger.warning("ReqKey validation failed closed: %s", exc)
292
+ await self._respond(
293
+ scope,
294
+ receive,
295
+ send,
296
+ status_code=503,
297
+ error="reqkey_unavailable",
298
+ )
299
+ return
300
+
301
+ if decision is not None and not decision.valid:
302
+ status_code, error = self._denial(decision)
303
+ headers = {}
304
+ if decision.retry_after is not None:
305
+ headers["Retry-After"] = str(max(0, int(decision.retry_after)))
306
+ await self._respond(
307
+ scope,
308
+ receive,
309
+ send,
310
+ status_code=status_code,
311
+ error=error,
312
+ headers=headers,
313
+ )
314
+ return
315
+
316
+ if decision is not None:
317
+ state = scope.setdefault("state", {})
318
+ state["reqkey"] = decision
319
+ state["reqkey_request_id"] = decision.request_id
320
+
321
+ if self.mode == "validate":
322
+ await self._run_validation_only(
323
+ scope,
324
+ receive,
325
+ send,
326
+ decision=decision,
327
+ validation_time_ms=validation_time_ms,
328
+ )
329
+ return
330
+
331
+ await self._run_with_blocking_ingest(
332
+ request,
333
+ scope,
334
+ receive,
335
+ send,
336
+ decision=decision,
337
+ validation_time_ms=validation_time_ms,
338
+ )
339
+
340
+ async def _run_validation_only(
341
+ self,
342
+ scope: Scope,
343
+ receive: Receive,
344
+ send: Send,
345
+ *,
346
+ decision: VerificationResult | None,
347
+ validation_time_ms: float | None,
348
+ ) -> None:
349
+ async def send_with_headers(message: Message) -> None:
350
+ self._add_decision_headers(message, decision, validation_time_ms)
351
+ await send(message)
352
+
353
+ await self.app(scope, receive, send_with_headers)
354
+
355
+ async def _run_with_blocking_ingest(
356
+ self,
357
+ request: Request,
358
+ scope: Scope,
359
+ receive: Receive,
360
+ send: Send,
361
+ *,
362
+ decision: VerificationResult | None,
363
+ validation_time_ms: float | None,
364
+ ) -> None:
365
+ response_status = 500
366
+ response_headers = Headers()
367
+ body_capture = _ResponseBodyCapture(enabled=False, headers=response_headers)
368
+ ingestion_completed = False
369
+ started = time.perf_counter()
370
+
371
+ async def stream_response(message: Message) -> None:
372
+ nonlocal response_status, response_headers, body_capture, ingestion_completed
373
+ if message["type"] == "http.response.start":
374
+ response_status = message["status"]
375
+ response_headers = Headers(raw=list(message.get("headers", [])))
376
+ body_capture = _ResponseBodyCapture(
377
+ enabled=self.capture_response_body,
378
+ headers=response_headers,
379
+ )
380
+ self._add_decision_headers(message, decision, validation_time_ms)
381
+ await send(message)
382
+ return
383
+
384
+ if message["type"] != "http.response.body":
385
+ await send(message)
386
+ return
387
+
388
+ body_capture.feed(message.get("body", b""))
389
+ if message.get("more_body", False):
390
+ await send(message)
391
+ return
392
+
393
+ await self._ingest_safely(
394
+ request=request,
395
+ decision=decision,
396
+ response_status=response_status,
397
+ latency_ms=round((time.perf_counter() - started) * 1000),
398
+ response_headers=response_headers,
399
+ response_body=body_capture.text(),
400
+ )
401
+ ingestion_completed = True
402
+ await send(message)
403
+
404
+ try:
405
+ await self.app(scope, receive, stream_response)
406
+ except Exception:
407
+ if not ingestion_completed:
408
+ await self._ingest_safely(
409
+ request=request,
410
+ decision=decision,
411
+ response_status=500,
412
+ latency_ms=round((time.perf_counter() - started) * 1000),
413
+ response_headers=response_headers,
414
+ response_body=body_capture.text(),
415
+ )
416
+ raise
417
+
418
+ if not ingestion_completed:
419
+ await self._ingest_safely(
420
+ request=request,
421
+ decision=decision,
422
+ response_status=response_status,
423
+ latency_ms=round((time.perf_counter() - started) * 1000),
424
+ response_headers=response_headers,
425
+ response_body=body_capture.text(),
426
+ )
427
+
428
+ async def _applies_to(self, request: Request) -> bool:
429
+ if request.method.upper() in self.skip_methods:
430
+ return False
431
+ if any(self._path_matches(request.url.path, pattern) for pattern in self.exclude_paths):
432
+ return False
433
+ if self.should_protect is None:
434
+ return True
435
+ result = self.should_protect(request)
436
+ return bool(await result) if inspect.isawaitable(result) else bool(result)
437
+
438
+ @staticmethod
439
+ def _path_matches(path: str, pattern: str) -> bool:
440
+ if pattern.endswith("*"):
441
+ return path.startswith(pattern[:-1])
442
+ return path == pattern
443
+
444
+ async def _consumer_key(self, request: Request) -> str | None:
445
+ if self.get_consumer_key is not None:
446
+ raw = self.get_consumer_key(request)
447
+ value = await raw if inspect.isawaitable(raw) else raw
448
+ elif self.key_location == "header":
449
+ value = request.headers.get(self.key_name)
450
+ elif self.key_location == "query":
451
+ value = request.query_params.get(self.key_name)
452
+ else:
453
+ value = request.cookies.get(self.key_name)
454
+
455
+ if value is None or not value.strip():
456
+ return None
457
+ value = value.strip()
458
+ if self.key_scheme == "bearer":
459
+ scheme, separator, credential = value.partition(" ")
460
+ if not separator or scheme.lower() != "bearer" or not credential.strip():
461
+ return None
462
+ return credential.strip()
463
+ return value
464
+
465
+ async def _credit_cost(self, request: Request) -> int:
466
+ raw = self.credits(request) if callable(self.credits) else self.credits
467
+ value = await raw if inspect.isawaitable(raw) else raw
468
+ if isinstance(value, bool) or not isinstance(value, int) or value < 0:
469
+ raise ReqKeyConfigurationError(
470
+ "The credits resolver must return a non-negative integer."
471
+ )
472
+ return value
473
+
474
+ async def _resource_path(self, request: Request) -> str:
475
+ if self.path_resolver is None:
476
+ return request.url.path
477
+ raw = self.path_resolver(request)
478
+ value = await raw if inspect.isawaitable(raw) else raw
479
+ if not value:
480
+ raise ReqKeyConfigurationError("The path resolver returned an empty path.")
481
+ return value
482
+
483
+ async def _resolved_request_id(
484
+ self,
485
+ request: Request,
486
+ decision: VerificationResult | None,
487
+ ) -> str | None:
488
+ if decision is not None:
489
+ return decision.request_id
490
+ if self.request_id_resolver is not None:
491
+ raw = self.request_id_resolver(request)
492
+ value = await raw if inspect.isawaitable(raw) else raw
493
+ return value or None
494
+ state_request_id = getattr(request.state, "reqkey_request_id", None)
495
+ return state_request_id if isinstance(state_request_id, str) else None
496
+
497
+ async def _ingest_safely(
498
+ self,
499
+ *,
500
+ request: Request,
501
+ decision: VerificationResult | None,
502
+ response_status: int,
503
+ latency_ms: int,
504
+ response_headers: Headers,
505
+ response_body: str | None,
506
+ ) -> None:
507
+ resource_path = await self._resource_path(request)
508
+ query_params, query = self._captured_query(request)
509
+ full_path = resource_path + (f"?{query}" if query else "")
510
+ request_headers = (
511
+ self._filtered_headers(request.headers)
512
+ if self.capture_request_headers
513
+ else None
514
+ )
515
+ captured_response_headers = (
516
+ self._filtered_headers(response_headers)
517
+ if self.capture_response_headers
518
+ else None
519
+ )
520
+ client_ip = (
521
+ request.client.host
522
+ if self.capture_client_ip and request.client is not None
523
+ else None
524
+ )
525
+ user_agent = request.headers.get("user-agent") if self.capture_user_agent else None
526
+
527
+ try:
528
+ await self.client.ingest(
529
+ await self._resolved_request_id(request, decision),
530
+ api_id=self.api_id,
531
+ method=request.method,
532
+ endpoint=resource_path,
533
+ path=full_path,
534
+ status_code=response_status,
535
+ latency_ms=latency_ms,
536
+ client_ip=client_ip,
537
+ user_agent=user_agent,
538
+ query_params=query_params,
539
+ request_headers=request_headers,
540
+ response_headers=captured_response_headers,
541
+ response_body=response_body,
542
+ )
543
+ except ReqKeyError as exc:
544
+ logger.warning("ReqKey analytics ingestion failed: %s", exc)
545
+
546
+ def _captured_query(self, request: Request) -> tuple[dict[str, str] | None, str]:
547
+ if not self.capture_query_params:
548
+ return None, ""
549
+ pairs = [
550
+ (key, value)
551
+ for key, value in request.query_params.multi_items()
552
+ if not (self.key_location == "query" and key == self.key_name)
553
+ ]
554
+ return dict(pairs), urlencode(pairs)
555
+
556
+ def _filtered_headers(self, headers: Mapping[str, str]) -> dict[str, str]:
557
+ return {
558
+ key: value
559
+ for key, value in headers.items()
560
+ if key.lower() not in self.excluded_headers
561
+ }
562
+
563
+ @staticmethod
564
+ def _add_decision_headers(
565
+ message: Message,
566
+ decision: VerificationResult | None,
567
+ validation_time_ms: float | None,
568
+ ) -> None:
569
+ if message["type"] != "http.response.start":
570
+ return
571
+ headers = MutableHeaders(scope=message)
572
+ if decision is not None:
573
+ if decision.request_id:
574
+ headers["X-ReqKey-Request-ID"] = decision.request_id
575
+ if decision.credits_limit is not None:
576
+ headers["X-ReqKey-Credits-Limit"] = str(decision.credits_limit)
577
+ if decision.credits_remaining is not None:
578
+ headers["X-ReqKey-Credits-Remaining"] = str(decision.credits_remaining)
579
+ if validation_time_ms is not None:
580
+ headers["X-ReqKey-Validation-Time-Ms"] = f"{validation_time_ms:.3f}"
581
+
582
+ @staticmethod
583
+ def _denial(decision: VerificationResult) -> tuple[int, str]:
584
+ if decision.reason is VerificationReason.INSUFFICIENT_CREDITS:
585
+ return 402, "insufficient_credits"
586
+ if decision.reason is VerificationReason.RATE_LIMITED:
587
+ return 429, "rate_limited"
588
+ if decision.reason is VerificationReason.FORBIDDEN:
589
+ return 403, "access_denied"
590
+ return 401, "invalid_api_key"
591
+
592
+ async def _respond(
593
+ self,
594
+ scope: Scope,
595
+ receive: Receive,
596
+ send: Send,
597
+ *,
598
+ status_code: int,
599
+ error: str,
600
+ headers: Mapping[str, str] | None = None,
601
+ ) -> None:
602
+ response = JSONResponse(
603
+ {"error": error, "message": self.error_messages[error]},
604
+ status_code=status_code,
605
+ headers=dict(headers or {}),
606
+ )
607
+ await response(scope, receive, send)
608
+
609
+
610
+ ReqKeyASGIMiddleware = ReqKeyMiddleware
611
+
612
+ __all__ = ["ReqKeyASGIMiddleware", "ReqKeyMiddleware"]