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/models.py ADDED
@@ -0,0 +1,45 @@
1
+ """Public response models for the ReqKey SDK."""
2
+
3
+ from __future__ import annotations
4
+
5
+ from dataclasses import dataclass, field
6
+ from enum import Enum
7
+ from typing import Any
8
+
9
+
10
+ class VerificationReason(str, Enum):
11
+ """Stable SDK-level categories for a key-validation decision."""
12
+
13
+ VALID = "valid"
14
+ INVALID_KEY = "invalid_key"
15
+ INSUFFICIENT_CREDITS = "insufficient_credits"
16
+ FORBIDDEN = "forbidden"
17
+ RATE_LIMITED = "rate_limited"
18
+ DENIED = "denied"
19
+
20
+
21
+ @dataclass(frozen=True, slots=True)
22
+ class VerificationResult:
23
+ """The result of a `/key/validate` decision."""
24
+
25
+ valid: bool
26
+ reason: VerificationReason
27
+ status_code: int
28
+ request_id: str | None = None
29
+ message: str | None = None
30
+ api_id: str | None = None
31
+ api_name: str | None = None
32
+ resource: str | None = None
33
+ credits_remaining: int | None = None
34
+ credits_limit: int | None = None
35
+ allowed_apis: tuple[str, ...] = ()
36
+ retry_after: float | None = None
37
+ rate_limit: dict[str, Any] | None = None
38
+ raw: dict[str, Any] = field(default_factory=dict, repr=False)
39
+
40
+ @property
41
+ def allowed(self) -> bool:
42
+ """Alias that reads naturally inside authorization code."""
43
+
44
+ return self.valid
45
+
reqkey/py.typed ADDED
@@ -0,0 +1 @@
1
+
reqkey/wsgi.py ADDED
@@ -0,0 +1,546 @@
1
+ """Generic WSGI middleware for Flask, Bottle, Pyramid, and other WSGI apps."""
2
+
3
+ from __future__ import annotations
4
+
5
+ import json
6
+ import logging
7
+ import time
8
+ import zlib
9
+ from collections.abc import Callable, Iterable, Iterator, Mapping
10
+ from datetime import datetime
11
+ from http.cookies import SimpleCookie
12
+ from types import TracebackType
13
+ from typing import Any, Protocol
14
+ from urllib.parse import parse_qsl, urlencode
15
+
16
+ from ._middleware import (
17
+ DEFAULT_ERROR_MESSAGES,
18
+ FailureMode,
19
+ KeyLocation,
20
+ KeyScheme,
21
+ Mode,
22
+ decision_headers,
23
+ denial,
24
+ excluded_header_names,
25
+ extract_credential,
26
+ filtered_headers,
27
+ path_matches,
28
+ validate_credit_cost,
29
+ validate_middleware_options,
30
+ )
31
+ from .client import DEFAULT_BASE_URL, DEFAULT_TIMEOUT_SECONDS, ReqKey
32
+ from .exceptions import ReqKeyConfigurationError, ReqKeyError
33
+ from .models import VerificationResult
34
+
35
+ logger = logging.getLogger("reqkey")
36
+
37
+ MAX_RESPONSE_BODY_CHARACTERS = 1000
38
+ _MAX_CAPTURE_BYTES = MAX_RESPONSE_BODY_CHARACTERS * 4 + 4
39
+
40
+ WSGIEnvironment = dict[str, Any]
41
+ ResponseHeaders = list[tuple[str, str]]
42
+ ExcInfo = tuple[type[BaseException], BaseException, TracebackType]
43
+
44
+
45
+ class StartResponse(Protocol):
46
+ def __call__(
47
+ self,
48
+ status: str,
49
+ response_headers: ResponseHeaders,
50
+ exc_info: ExcInfo | None = None,
51
+ ) -> Callable[[bytes], object]: ...
52
+
53
+
54
+ class WSGIApplication(Protocol):
55
+ def __call__(
56
+ self,
57
+ environ: WSGIEnvironment,
58
+ start_response: StartResponse,
59
+ ) -> Iterable[bytes]: ...
60
+
61
+
62
+ class ReqKeyLike(Protocol):
63
+ def verify(
64
+ self,
65
+ key: str,
66
+ *,
67
+ api_id: str | None = None,
68
+ credits: int = 1,
69
+ resource: str | None = None,
70
+ ) -> VerificationResult: ...
71
+
72
+ def ingest(
73
+ self,
74
+ request_id: str | None = None,
75
+ *,
76
+ api_id: str | None = None,
77
+ method: str | None = None,
78
+ endpoint: str | None = None,
79
+ path: str | None = None,
80
+ status_code: int | None = None,
81
+ latency_ms: int | None = None,
82
+ client_ip: str | None = None,
83
+ user_agent: str | None = None,
84
+ user_id: str | None = None,
85
+ query_params: Mapping[str, Any] | None = None,
86
+ request_headers: Mapping[str, str] | None = None,
87
+ response_headers: Mapping[str, str] | None = None,
88
+ request_body: str | None = None,
89
+ response_body: str | None = None,
90
+ timestamp: datetime | str | None = None,
91
+ ) -> None: ...
92
+
93
+
94
+ CreditsResolver = Callable[[WSGIEnvironment], int]
95
+ ProtectionResolver = Callable[[WSGIEnvironment], bool]
96
+ ConsumerKeyResolver = Callable[[WSGIEnvironment], str | None]
97
+ RequestIdResolver = Callable[[WSGIEnvironment], str | None]
98
+ PathResolver = Callable[[WSGIEnvironment], str]
99
+
100
+
101
+ class _ResponseBodyCapture:
102
+ def __init__(self, *, enabled: bool, headers: Mapping[str, str]) -> None:
103
+ content_type = headers.get("content-type", "").lower()
104
+ self.enabled = enabled and (
105
+ not content_type
106
+ or any(
107
+ marker in content_type
108
+ for marker in (
109
+ "json",
110
+ "text/",
111
+ "xml",
112
+ "javascript",
113
+ "x-www-form-urlencoded",
114
+ )
115
+ )
116
+ )
117
+ self._body = bytearray()
118
+ self._decompressor = (
119
+ zlib.decompressobj(16 + zlib.MAX_WBITS)
120
+ if self.enabled and headers.get("content-encoding", "").lower() == "gzip"
121
+ else None
122
+ )
123
+
124
+ def feed(self, chunk: bytes) -> None:
125
+ if not self.enabled or not chunk or len(self._body) >= _MAX_CAPTURE_BYTES:
126
+ return
127
+ remaining = _MAX_CAPTURE_BYTES - len(self._body)
128
+ if self._decompressor is None:
129
+ self._body.extend(chunk[:remaining])
130
+ return
131
+ try:
132
+ self._body.extend(self._decompressor.decompress(chunk, remaining))
133
+ except zlib.error:
134
+ self.enabled = False
135
+ self._body.clear()
136
+ logger.debug("Could not decompress a gzip response for ReqKey analytics.")
137
+
138
+ def text(self) -> str | None:
139
+ if not self.enabled:
140
+ return None
141
+ if self._decompressor is not None and len(self._body) < _MAX_CAPTURE_BYTES:
142
+ try:
143
+ remaining = _MAX_CAPTURE_BYTES - len(self._body)
144
+ self._body.extend(self._decompressor.flush(remaining))
145
+ except zlib.error:
146
+ logger.debug("Could not finish gzip response capture for ReqKey analytics.")
147
+ return self._body.decode("utf-8", errors="replace")[:MAX_RESPONSE_BODY_CHARACTERS]
148
+
149
+
150
+ class _ResponseState:
151
+ def __init__(self) -> None:
152
+ self.status_code = 500
153
+ self.headers: dict[str, str] = {}
154
+ self.body = _ResponseBodyCapture(enabled=False, headers={})
155
+
156
+
157
+ class ReqKeyMiddleware:
158
+ """Pure WSGI middleware for incoming API validation and analytics."""
159
+
160
+ def __init__(
161
+ self,
162
+ app: WSGIApplication,
163
+ *,
164
+ api_id: str,
165
+ project_key: str | None = None,
166
+ root_key: str | None = None,
167
+ client: ReqKeyLike | None = None,
168
+ base_url: str = DEFAULT_BASE_URL,
169
+ timeout: float = DEFAULT_TIMEOUT_SECONDS,
170
+ mode: Mode = "both",
171
+ enabled: bool = True,
172
+ key_location: KeyLocation = "header",
173
+ key_name: str = "X-API-Key",
174
+ key_scheme: KeyScheme = "raw",
175
+ get_consumer_key: ConsumerKeyResolver | None = None,
176
+ credits: int | CreditsResolver = 1,
177
+ exclude_paths: Iterable[str] = (),
178
+ skip_methods: Iterable[str] = ("OPTIONS",),
179
+ should_protect: ProtectionResolver | None = None,
180
+ request_id_resolver: RequestIdResolver | None = None,
181
+ path_resolver: PathResolver | None = None,
182
+ error_messages: Mapping[str, str] | None = None,
183
+ capture_query_params: bool = False,
184
+ capture_request_headers: bool = False,
185
+ capture_response_headers: bool = False,
186
+ capture_response_body: bool = False,
187
+ capture_client_ip: bool = False,
188
+ capture_user_agent: bool = True,
189
+ excluded_headers: Iterable[str] = (),
190
+ failure_mode: FailureMode = "closed",
191
+ ) -> None:
192
+ validate_middleware_options(
193
+ api_id=api_id,
194
+ mode=mode,
195
+ key_location=key_location,
196
+ key_scheme=key_scheme,
197
+ failure_mode=failure_mode,
198
+ credits=credits,
199
+ )
200
+ if client is not None and (project_key is not None or root_key is not None):
201
+ raise ReqKeyConfigurationError("Pass client or project_key/root_key, not both.")
202
+
203
+ self.app = app
204
+ self.api_id = api_id
205
+ self.client = client or ReqKey(
206
+ project_key=project_key,
207
+ root_key=root_key,
208
+ base_url=base_url,
209
+ timeout=timeout,
210
+ )
211
+ self._owns_client = client is None
212
+ self.mode = mode
213
+ self.enabled = enabled
214
+ self.key_location = key_location
215
+ self.key_name = key_name
216
+ self.key_scheme = key_scheme
217
+ self.get_consumer_key = get_consumer_key
218
+ self.credits = credits
219
+ self.exclude_paths = tuple(exclude_paths)
220
+ self.skip_methods = frozenset(method.upper() for method in skip_methods)
221
+ self.should_protect = should_protect
222
+ self.request_id_resolver = request_id_resolver
223
+ self.path_resolver = path_resolver
224
+ self.error_messages = {**DEFAULT_ERROR_MESSAGES, **dict(error_messages or {})}
225
+ self.capture_query_params = capture_query_params
226
+ self.capture_request_headers = capture_request_headers
227
+ self.capture_response_headers = capture_response_headers
228
+ self.capture_response_body = capture_response_body
229
+ self.capture_client_ip = capture_client_ip
230
+ self.capture_user_agent = capture_user_agent
231
+ self.excluded_headers = excluded_header_names(key_name, excluded_headers)
232
+ self.failure_mode = failure_mode
233
+
234
+ def __call__(
235
+ self,
236
+ environ: WSGIEnvironment,
237
+ start_response: StartResponse,
238
+ ) -> Iterable[bytes]:
239
+ if not self.enabled or not self._applies_to(environ):
240
+ return self.app(environ, start_response)
241
+
242
+ decision: VerificationResult | None = None
243
+ validation_time_ms: float | None = None
244
+
245
+ if self.mode in {"validate", "both"}:
246
+ consumer_key = self._consumer_key(environ)
247
+ if consumer_key is None:
248
+ return self._respond(start_response, 401, "missing_api_key")
249
+
250
+ validation_started = time.perf_counter()
251
+ try:
252
+ decision = self.client.verify(
253
+ consumer_key,
254
+ api_id=self.api_id,
255
+ credits=self._credit_cost(environ),
256
+ resource=self._resource_path(environ),
257
+ )
258
+ validation_time_ms = (time.perf_counter() - validation_started) * 1000
259
+ except ReqKeyError as exc:
260
+ validation_time_ms = (time.perf_counter() - validation_started) * 1000
261
+ if self.failure_mode == "open":
262
+ environ["reqkey.error"] = exc
263
+ else:
264
+ logger.warning("ReqKey validation failed closed: %s", exc)
265
+ return self._respond(start_response, 503, "reqkey_unavailable")
266
+
267
+ if decision is not None and not decision.valid:
268
+ status_code, error = denial(decision)
269
+ headers = {}
270
+ if decision.retry_after is not None:
271
+ headers["Retry-After"] = str(max(0, int(decision.retry_after)))
272
+ return self._respond(start_response, status_code, error, headers=headers)
273
+
274
+ if decision is not None:
275
+ environ["reqkey.decision"] = decision
276
+ environ["reqkey.request_id"] = decision.request_id
277
+
278
+ headers = decision_headers(decision, validation_time_ms)
279
+ wrapped_start = self._start_response_with_headers(start_response, headers)
280
+ if self.mode == "validate":
281
+ return self.app(environ, wrapped_start)
282
+
283
+ return self._run_with_ingest(
284
+ environ,
285
+ wrapped_start,
286
+ decision=decision,
287
+ )
288
+
289
+ def _run_with_ingest(
290
+ self,
291
+ environ: WSGIEnvironment,
292
+ start_response: StartResponse,
293
+ *,
294
+ decision: VerificationResult | None,
295
+ ) -> Iterator[bytes]:
296
+ response = _ResponseState()
297
+ started = time.perf_counter()
298
+
299
+ def capture_start_response(
300
+ status: str,
301
+ response_headers: ResponseHeaders,
302
+ exc_info: ExcInfo | None = None,
303
+ ) -> Callable[[bytes], object]:
304
+ response.status_code = _status_code(status)
305
+ response.headers = {key.lower(): value for key, value in response_headers}
306
+ response.body = _ResponseBodyCapture(
307
+ enabled=self.capture_response_body,
308
+ headers=response.headers,
309
+ )
310
+ write = start_response(status, response_headers, exc_info)
311
+
312
+ def capture_write(data: bytes) -> object:
313
+ response.body.feed(data)
314
+ return write(data)
315
+
316
+ return capture_write
317
+
318
+ app_iterable: Iterable[bytes] | None = None
319
+ failed = False
320
+ try:
321
+ app_iterable = self.app(environ, capture_start_response)
322
+ for chunk in app_iterable:
323
+ if not isinstance(chunk, bytes):
324
+ raise TypeError("WSGI applications must yield bytes.")
325
+ response.body.feed(chunk)
326
+ yield chunk
327
+ except Exception:
328
+ failed = True
329
+ raise
330
+ finally:
331
+ if app_iterable is not None:
332
+ close = getattr(app_iterable, "close", None)
333
+ if callable(close):
334
+ close()
335
+ self._ingest_safely(
336
+ environ=environ,
337
+ decision=decision,
338
+ response_status=500 if failed else response.status_code,
339
+ latency_ms=round((time.perf_counter() - started) * 1000),
340
+ response_headers=response.headers,
341
+ response_body=response.body.text(),
342
+ )
343
+
344
+ def _applies_to(self, environ: WSGIEnvironment) -> bool:
345
+ method = str(environ.get("REQUEST_METHOD", "GET")).upper()
346
+ path = str(environ.get("PATH_INFO", "/")) or "/"
347
+ if method in self.skip_methods:
348
+ return False
349
+ if any(path_matches(path, pattern) for pattern in self.exclude_paths):
350
+ return False
351
+ return True if self.should_protect is None else bool(self.should_protect(environ))
352
+
353
+ def _consumer_key(self, environ: WSGIEnvironment) -> str | None:
354
+ if self.get_consumer_key is not None:
355
+ value = self.get_consumer_key(environ)
356
+ elif self.key_location == "header":
357
+ value = _request_header(environ, self.key_name)
358
+ elif self.key_location == "query":
359
+ values = dict(parse_qsl(str(environ.get("QUERY_STRING", "")), keep_blank_values=True))
360
+ value = values.get(self.key_name)
361
+ else:
362
+ cookie = SimpleCookie()
363
+ cookie.load(str(environ.get("HTTP_COOKIE", "")))
364
+ morsel = cookie.get(self.key_name)
365
+ value = morsel.value if morsel is not None else None
366
+ return extract_credential(value, self.key_scheme)
367
+
368
+ def _credit_cost(self, environ: WSGIEnvironment) -> int:
369
+ raw = self.credits(environ) if callable(self.credits) else self.credits
370
+ return validate_credit_cost(raw)
371
+
372
+ def _resource_path(self, environ: WSGIEnvironment) -> str:
373
+ if self.path_resolver is None:
374
+ return str(environ.get("PATH_INFO", "/")) or "/"
375
+ value = self.path_resolver(environ)
376
+ if not value:
377
+ raise ReqKeyConfigurationError("The path resolver returned an empty path.")
378
+ return value
379
+
380
+ def _resolved_request_id(
381
+ self,
382
+ environ: WSGIEnvironment,
383
+ decision: VerificationResult | None,
384
+ ) -> str | None:
385
+ if decision is not None:
386
+ return decision.request_id
387
+ if self.request_id_resolver is not None:
388
+ return self.request_id_resolver(environ) or None
389
+ value = environ.get("reqkey.request_id")
390
+ return value if isinstance(value, str) else None
391
+
392
+ def _ingest_safely(
393
+ self,
394
+ *,
395
+ environ: WSGIEnvironment,
396
+ decision: VerificationResult | None,
397
+ response_status: int,
398
+ latency_ms: int,
399
+ response_headers: Mapping[str, str],
400
+ response_body: str | None,
401
+ ) -> None:
402
+ resource_path = self._resource_path(environ)
403
+ query_params, query = self._captured_query(environ)
404
+ full_path = resource_path + (f"?{query}" if query else "")
405
+ request_headers = _request_headers(environ)
406
+
407
+ try:
408
+ self.client.ingest(
409
+ self._resolved_request_id(environ, decision),
410
+ api_id=self.api_id,
411
+ method=str(environ.get("REQUEST_METHOD", "GET")),
412
+ endpoint=resource_path,
413
+ path=full_path,
414
+ status_code=response_status,
415
+ latency_ms=latency_ms,
416
+ client_ip=(
417
+ str(environ.get("REMOTE_ADDR"))
418
+ if self.capture_client_ip and environ.get("REMOTE_ADDR") is not None
419
+ else None
420
+ ),
421
+ user_agent=(
422
+ request_headers.get("user-agent") if self.capture_user_agent else None
423
+ ),
424
+ query_params=query_params,
425
+ request_headers=(
426
+ filtered_headers(request_headers, self.excluded_headers)
427
+ if self.capture_request_headers
428
+ else None
429
+ ),
430
+ response_headers=(
431
+ filtered_headers(response_headers, self.excluded_headers)
432
+ if self.capture_response_headers
433
+ else None
434
+ ),
435
+ response_body=response_body,
436
+ )
437
+ except ReqKeyError as exc:
438
+ logger.warning("ReqKey analytics ingestion failed: %s", exc)
439
+
440
+ def _captured_query(
441
+ self,
442
+ environ: WSGIEnvironment,
443
+ ) -> tuple[dict[str, str] | None, str]:
444
+ if not self.capture_query_params:
445
+ return None, ""
446
+ pairs = [
447
+ (key, value)
448
+ for key, value in parse_qsl(
449
+ str(environ.get("QUERY_STRING", "")),
450
+ keep_blank_values=True,
451
+ )
452
+ if not (self.key_location == "query" and key == self.key_name)
453
+ ]
454
+ return dict(pairs), urlencode(pairs)
455
+
456
+ @staticmethod
457
+ def _start_response_with_headers(
458
+ start_response: StartResponse,
459
+ added_headers: Mapping[str, str],
460
+ ) -> StartResponse:
461
+ def wrapped(
462
+ status: str,
463
+ response_headers: ResponseHeaders,
464
+ exc_info: ExcInfo | None = None,
465
+ ) -> Callable[[bytes], object]:
466
+ headers = _merge_headers(response_headers, added_headers)
467
+ return start_response(status, headers, exc_info)
468
+
469
+ return wrapped
470
+
471
+ def _respond(
472
+ self,
473
+ start_response: StartResponse,
474
+ status_code: int,
475
+ error: str,
476
+ *,
477
+ headers: Mapping[str, str] | None = None,
478
+ ) -> Iterable[bytes]:
479
+ body = json.dumps(
480
+ {"error": error, "message": self.error_messages[error]},
481
+ separators=(",", ":"),
482
+ ).encode()
483
+ response_headers: ResponseHeaders = [
484
+ ("Content-Type", "application/json"),
485
+ ("Content-Length", str(len(body))),
486
+ ]
487
+ response_headers.extend((key, value) for key, value in (headers or {}).items())
488
+ start_response(f"{status_code} {_reason_phrase(status_code)}", response_headers)
489
+ return [body]
490
+
491
+ def close(self) -> None:
492
+ if self._owns_client and isinstance(self.client, ReqKey):
493
+ self.client.close()
494
+
495
+
496
+ ReqKeyWSGIMiddleware = ReqKeyMiddleware
497
+
498
+
499
+ def _request_header(environ: WSGIEnvironment, name: str) -> str | None:
500
+ normalized = name.upper().replace("-", "_")
501
+ key = normalized if normalized in {"CONTENT_TYPE", "CONTENT_LENGTH"} else f"HTTP_{normalized}"
502
+ value = environ.get(key)
503
+ return str(value) if value is not None else None
504
+
505
+
506
+ def _request_headers(environ: WSGIEnvironment) -> dict[str, str]:
507
+ headers: dict[str, str] = {}
508
+ for key, value in environ.items():
509
+ if key.startswith("HTTP_"):
510
+ name = key[5:].replace("_", "-").lower()
511
+ elif key in {"CONTENT_TYPE", "CONTENT_LENGTH"}:
512
+ name = key.replace("_", "-").lower()
513
+ else:
514
+ continue
515
+ headers[name] = str(value)
516
+ return headers
517
+
518
+
519
+ def _status_code(status: str) -> int:
520
+ try:
521
+ return int(status.partition(" ")[0])
522
+ except ValueError:
523
+ return 500
524
+
525
+
526
+ def _merge_headers(
527
+ headers: ResponseHeaders,
528
+ additions: Mapping[str, str],
529
+ ) -> ResponseHeaders:
530
+ lowered = {key.lower() for key in additions}
531
+ merged = [(key, value) for key, value in headers if key.lower() not in lowered]
532
+ merged.extend(additions.items())
533
+ return merged
534
+
535
+
536
+ def _reason_phrase(status_code: int) -> str:
537
+ return {
538
+ 401: "Unauthorized",
539
+ 402: "Payment Required",
540
+ 403: "Forbidden",
541
+ 429: "Too Many Requests",
542
+ 503: "Service Unavailable",
543
+ }.get(status_code, "Error")
544
+
545
+
546
+ __all__ = ["ReqKeyMiddleware", "ReqKeyWSGIMiddleware"]