diff --git a/.github/RELEASE_NOTES_v0.8.0.md b/.github/RELEASE_NOTES_v0.8.0.md new file mode 100644 index 0000000..94ead4e --- /dev/null +++ b/.github/RELEASE_NOTES_v0.8.0.md @@ -0,0 +1,35 @@ +# OxyJWT 0.8.0 + +**Beta** performance release — cached RSA signing keys, borrowed (no longer cloned) key material, a lighter Python decode fast path, interned claim/header names, and skipped GIL release on the HMAC hot path. No intentional breaking changes to the public `__all__` API. + +## Highlights + +### Performance + +- **Cached RSA signing key** — `EncodingKey.from_rsa_pem` parses the DER key into an `aws_lc_rs::RsaKeyPair` once, at construction time; RS256 `encode` dropped from ~526 µs to ~201 µs (2.6× faster) with a pre-built key +- **No more per-call key cloning** — `EncodingKey` / `DecodingKey` are `frozen` pyclasses now; native `encode`/`decode` borrow the key material instead of cloning it +- **Lighter Python decode fast path** — inlined argument checks, and `exp` is only re-checked in Python when it isn't a plain `int`: wrapper overhead dropped from ~0.85 µs to ~0.57 µs (with `iat`) or ~0.37 µs (no time claims) +- **Interned claim/header names** (`exp`, `iat`, `nbf`, `sub`, `aud`, `iss`, `jti`, `alg`, `typ`, `kid`) — native `decode` on an 8-claim payload dropped from ~2.05 µs to ~1.92 µs +- **HMAC no longer releases the GIL** — HS256 decode throughput under 8-thread contention went from ~410k to ~790k decodes/sec; RSA/EC/EdDSA are unaffected and keep releasing the GIL +- Removed several redundant allocations/re-parses on secondary paths (`get_unverified_header`, `decode_unverified`, the unverified branch of `decode_complete`, `encode_json` token assembly) + +### Behaviour change + +- `decode_unverified` now accepts a header that is valid JSON but not a recognized `alg` name (e.g. `{"alg": "made-up"}`), matching `get_unverified_header`'s own check. Unverified decode was never a security boundary. + +## Install + +```bash +pip install oxyjwt==0.8.0 +``` + +## Upgrade from 0.7.0 + +```bash +pip install -U oxyjwt +``` + +- No intentional breaking changes to public symbols. +- Successful encode/decode results are unchanged; only the `decode_unverified` edge case above differs. + +See the full [changelog](https://github.com/QueryaHub/OxyJWT/blob/main/CHANGELOG.md#080--2026-09-28). diff --git a/CHANGELOG.md b/CHANGELOG.md index 17021a6..ae5b1d2 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -9,6 +9,128 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 (No changes yet.) +## [0.8.0] — 2026-09-28 + +### Performance + +- **RSA/RSA-PSS `encode` no longer re-parses the private key on every call.** + `EncodingKey.from_rsa_pem` now parses the DER-encoded key into an + `aws_lc_rs::signature::RsaKeyPair` once, at construction time, and `encode` + signs through that cached key directly. Previously `jsonwebtoken::crypto::sign` + ran `RsaKeyPair::from_der` (including full RSA key validation) on every `encode` + call, which dominated the cost of signing. Measured with a pre-built + `EncodingKey` and a 2048-bit key: RS256 `encode` dropped from ~526 µs to + ~201 µs per call (~2.6× faster), now within ~5% of `cryptography`'s raw RSA + sign with an equivalent pre-parsed key. `decode`, and `encode`/`decode` for + HMAC, EC and EdDSA, are unaffected — those paths were already close to their + theoretical floor. Only active with the default `aws_lc_rs` crypto backend; + the `rust_crypto` feature (used for the Linux aarch64 wheel) is unchanged. + A malformed RSA private key is now rejected by `EncodingKey.from_rsa_pem` + itself instead of by the first `encode` call. (#120) +- **`EncodingKey` / `DecodingKey` no longer cloned on every `encode`/`decode` + call.** Both pyclasses are now `frozen`, and the native `encode`, `decode` + and `decode_complete` entry points borrow the underlying key material + straight out of the Python object instead of cloning it (an owned copy of + the DER/secret bytes, cloned again by `jsonwebtoken`'s signer/verifier + factory) on every call. HMAC secrets passed as raw `str`/`bytes` are + unaffected — there is no persistent key object to borrow from in that case. + No behavioural change. (#121) +- **Less Python-side work on the plain `decode`/`decode_complete` fast path.** + The `_is_plain_decode` argument check and the RFC 7797 `detached_payload` + pre-check are now inlined into `decode`/`decode_complete` instead of going + through a 9-argument and a keyword-argument function call. `exp` is only + re-checked in Python when it is not a plain `int`: Rust already enforces + `exp > now` with an integer clock on this path, which is the exact same + predicate for an integer `exp`, so re-running it only cost a `time.time()` + call with no behavioural difference; a `float` `exp` still gets the Python + recheck, since Rust rounds a fractional value to the nearest second where + PyJWT (and this check, for parity) truncates it, and the two can disagree + right at the boundary. Measured (HS256, prebuilt token, `int` `exp`): + wrapper overhead over the native call dropped from ~0.85 microseconds to + ~0.57 microseconds (~33% less); with no time claims at all, ~0.37 + microseconds (~57% less). No behavioural change. (#122) +- **Common claim/header names are interned instead of allocated per + decode.** `json_to_bound` (used by every `decode` and `decode_complete` + call) used to allocate a fresh `PyString` for every dict key, including + the same standard names — `exp`, `iat`, `nbf`, `sub`, `aud`, `iss`, `jti`, + `alg`, `typ`, `kid` — on every single call. Those ten keys now come from + `pyo3::intern!`, a per-key cache that both skips the repeated allocation + and interns the string in CPython's own intern table, so a later + `payload.get("exp")` on the Python side can hit the identity-comparison + fast path for dict lookups. Any other key still gets an ordinary, + uninterned `PyString`, unchanged from before. Measured on an 8-claim + payload: native `decode` dropped from ~2.05 µs to ~1.92 µs (~6% less). No + behavioural change. (#123) +- **HMAC `encode`/`decode` no longer release the GIL.** `encode`, + `encode_json`, `decode` and `decode_complete` used to call `py.detach` + unconditionally, releasing and reacquiring the GIL around every native + call. For HS256/384/512 the signing/verification itself takes roughly a + microsecond, so under thread contention the release-and-reacquire cycle + cost as much as the operation, or more: measured with 8 threads + continuously decoding the same HS256 token, throughput went from ~410k to + ~790k decodes/sec (~1.9× more) once the release was skipped, and + single-threaded decode dropped from ~1.18 µs to ~1.15 µs. RSA, EC and + EdDSA are unaffected — the check is on the resolved algorithm (`encode`) + or on the caller's allow-list (`decode`, checked before the algorithm in + the token header is known: an HMAC-only allow-list already guarantees the + verified algorithm is HMAC too), and those algorithms still release the + GIL, confirmed to keep scaling with threads (RS256 decode: ~67k/sec on 1 + thread, ~339k/sec on 8). No behavioural change; safe on free-threaded + Python 3.13t/3.14t since the HMAC path never calls back into Python + either way, so holding the GIL throughout is never a hazard, only a + choice not to release it. (#124) +- **Removed several small redundant allocations and re-parses on secondary + decode/encode paths.** None of these are on the main verified `decode` + hot path (already addressed by earlier entries in this section); each is + a modest, measured win on its own function: + - `get_unverified_header` no longer copies the token into an owned + `String` before `py.detach` (a borrowed `&str` is `Ungil` already; no + `'static` bound requires the copy) and no longer runs its own + `split_compact_segments` pre-check before `parse_compact_header_json` + runs the exact same split internally. ~384 ns → ~352 ns. + - `decode_unverified` used `jsonwebtoken::dangerous::insecure_decode`, + which fully deserializes the header into `jsonwebtoken`'s typed + `Header` struct even though only `.claims` was ever read, and + re-implements its own lenient segment split (silently misparsing a + token with extra `.`s instead of rejecting it, unlike our own + `split_compact_segments`) -- on top of the same redundant pre-check as + `get_unverified_header`. Replaced with a single-pass helper that + reuses our own strict split and only parses the header far enough to + confirm it is a JSON object (matching `get_unverified_header`'s own + check) before discarding it. ~681 ns → ~595 ns. + - `decode_complete`'s unverified path (`jws_parse_compact`) computed and + returned a `header.payload` "signing input" byte string that its only + Python caller immediately discarded; it no longer computes it at all. + ~787 ns → ~666 ns for the `encode_json` counterpart exercised by the + same benchmark payload (the byte-building change below); the + `decode_complete(verify_signature=False)` path itself is dominated by + Python-side claim validation, so the native saving there is smaller + (~2580 ns → ~2500 ns end to end). + - `decode_complete` (verified path) decoded the signature segment's + base64 a second time via a separate `extract_signature_bytes` call, + which re-split the *entire* token from scratch to reach it. It now + decodes the signature once, inline, right where the token is already + split for verification -- removing the redundant full re-split; + `jsonwebtoken`'s `crypto::verify` still does its own internal base64 + decode of just the (small, bounded) signature segment, since it has + no public entry point that accepts already-decoded signature bytes. + - `encode_json` (and the RSA fast path shared with `encode`) built the + `header.payload.signature` token through a chain of `Engine::encode` + calls into throwaway `String`s and two `format!`s, each copying + everything built so far into a new allocation. It now encodes header + and payload directly into one pre-sized `String` and appends the + signature to the same buffer, so the whole token is built with (at + most) one buffer growth instead of several full copies. ~787 ns → + ~666 ns. + + No behavioural change, other than `decode_unverified` becoming slightly + *more* lenient in one narrow, untested edge case: a token whose header is + valid JSON but not a recognized `alg` name (e.g. `{"alg": "made-up"}`) + now decodes instead of raising, aligning it with `get_unverified_header` + -- which already only required the header to be a JSON object -- rather + than `jsonwebtoken`'s stricter typed deserialization, which no other + method in this library performs for an *unverified* decode. (#125) + ## [0.7.0] — 2026-08-26 Performance release. Verified `decode` is about **2.3× faster** and `encode` about @@ -303,7 +425,8 @@ Initial alpha release. - Mixed algorithm families are rejected for one decode call. - In 0.1.0, `verify_signature=False` was rejected in `decode` (0.2.0 allows an explicit unverified path). -[Unreleased]: https://github.com/QueryaHub/OxyJWT/compare/v0.7.0...HEAD +[Unreleased]: https://github.com/QueryaHub/OxyJWT/compare/v0.8.0...HEAD +[0.8.0]: https://github.com/QueryaHub/OxyJWT/compare/v0.7.0...v0.8.0 [0.7.0]: https://github.com/QueryaHub/OxyJWT/compare/v0.6.0...v0.7.0 [0.6.0]: https://github.com/QueryaHub/OxyJWT/compare/v0.5.0...v0.6.0 [0.5.0]: https://github.com/QueryaHub/OxyJWT/compare/v0.4.0...v0.5.0 diff --git a/RELEASING.md b/RELEASING.md index cff41e7..d6fcd39 100644 --- a/RELEASING.md +++ b/RELEASING.md @@ -1,6 +1,6 @@ # Releasing OxyJWT -This checklist is for maintainers publishing **0.7.x** (and later) to PyPI via the GitHub Actions [Release workflow](.github/workflows/release.yml). +This checklist is for maintainers publishing **0.8.x** (and later) to PyPI via the GitHub Actions [Release workflow](.github/workflows/release.yml). ## Before tagging @@ -41,16 +41,16 @@ This checklist is for maintainers publishing **0.7.x** (and later) to PyPI via t ## Publish 1. Commit all release-prep changes on `main`. -2. Merge `dev` → `main`, then create and push an annotated tag (example for **0.7.0**): +2. Merge `dev` → `main`, then create and push an annotated tag (example for **0.8.0**): ```bash - git tag -a v0.7.0 -m "Release 0.7.0" - git push origin v0.7.0 + git tag -a v0.8.0 -m "Release 0.8.0" + git push origin v0.8.0 ``` 3. The **Release** workflow runs full [CI](.github/workflows/ci.yml) via `workflow_call`, then builds wheels (Linux x86_64/aarch64, macOS, Windows) + sdist and publishes to PyPI only if CI passes (requires the `pypi` environment and [Trusted Publishing](https://docs.pypi.org/trusted-publishers/)). -4. On GitHub, create a **Release** from the tag. Use [`.github/RELEASE_NOTES_v0.7.0.md`](.github/RELEASE_NOTES_v0.7.0.md) or the **0.7.0** section in `CHANGELOG.md` as the release notes body. +4. On GitHub, create a **Release** from the tag. Use [`.github/RELEASE_NOTES_v0.8.0.md`](.github/RELEASE_NOTES_v0.8.0.md) or the **0.8.0** section in `CHANGELOG.md` as the release notes body. ## After release diff --git a/docs-site/docs/changelog.md b/docs-site/docs/changelog.md index 056a41a..c2a8e3f 100644 --- a/docs-site/docs/changelog.md +++ b/docs-site/docs/changelog.md @@ -4,6 +4,78 @@ (No changes yet.) +## 0.8.0 — 2026-09-28 + +Performance release: cached RSA signing keys, borrowed (no longer cloned) key +material, a lighter Python decode fast path, interned claim/header names, and +skipped GIL release on the HMAC hot path. No intentional breaking changes to +the public `__all__` surface. The API remains pre-1.0 (Beta). See +[Versioning](versioning.md) and [`SECURITY.md` on GitHub](https://github.com/QueryaHub/OxyJWT/blob/main/SECURITY.md). + +### Upgrading from 0.7.0 + +```bash +pip install -U oxyjwt +``` + +- No intentional breaking changes to the public `__all__` surface. +- One edge-case behaviour change: `decode_unverified` now accepts a header + that is valid JSON but not a recognized `alg` name (for example + `{"alg": "made-up"}`), matching `get_unverified_header`'s own, looser + check. Unverified decode was never a security boundary, so this only + affects inspecting an untrusted token's claims without checking its + signature. + +| Operation | Before | After | Change | +| --- | --- | --- | --- | +| RS256 `encode` (pre-built key) | 526 µs | 201 µs | 2.6× faster | +| `encode`/`decode` (pre-built key, clone removed) | — | — | no measurable regression, fewer allocations | +| `decode` wrapper overhead, `int` `exp` + `iat` | 0.85 µs | 0.57 µs | 33% less | +| `decode` wrapper overhead, no time claims | 0.85 µs | 0.37 µs | 57% less | +| native `decode`, 8-claim payload (key interning) | 2.05 µs | 1.92 µs | 6% less | +| HS256 `decode`, 8 threads (no GIL release) | ~410k/s | ~790k/s | 1.9× more | +| `get_unverified_header` | 384 ns | 352 ns | 8% less | +| `decode_unverified` | 681 ns | 595 ns | 13% less | +| `encode_json` | 787 ns | 666 ns | 15% less | + +### Changed + +- **RSA/RSA-PSS `encode` caches the parsed private key.** `EncodingKey.from_rsa_pem` + parses the DER-encoded key into an `aws_lc_rs` `RsaKeyPair` once, at + construction time, instead of `jsonwebtoken::crypto::sign` re-parsing (and + re-validating) it on every `encode` call. A malformed RSA private key is + now rejected by `EncodingKey.from_rsa_pem` itself instead of by the first + `encode` call. Only active with the default `aws_lc_rs` crypto backend; + the `rust_crypto` feature (Linux aarch64 wheel) is unchanged. +- **`EncodingKey` / `DecodingKey` are no longer cloned on every `encode` / + `decode` call.** Both pyclasses are now `frozen`, so the native entry + points borrow the underlying key material directly instead of cloning it. + Raw `str` / `bytes` HMAC secrets are unaffected. +- **Less Python-side work on the plain `decode` / `decode_complete` fast + path.** The argument checks that used to go through two extra function + calls are now inlined, and `exp` is only re-checked in Python when it is + not a plain `int` (Rust's own integer-clock check already covers that + case exactly; a `float` `exp` still gets the Python recheck for + boundary-rounding parity with PyJWT). +- **Common claim / header names are interned** (`exp`, `iat`, `nbf`, `sub`, + `aud`, `iss`, `jti`, `alg`, `typ`, `kid`) instead of allocated fresh on + every decode. +- **HMAC `encode` / `decode` no longer release the GIL.** HS256/384/512 + sign/verify is fast enough that the release/reacquire cycle cost more + than the operation itself under thread contention. RSA, EC and EdDSA are + unaffected and keep releasing the GIL. +- Removed several redundant allocations and re-parses on secondary + decode/encode paths (`get_unverified_header`, `decode_unverified`, the + unverified branch of `decode_complete`, and `encode_json`'s token + assembly). + +### Behaviour change + +- `decode_unverified` now accepts a header that is valid JSON but is not a + recognized `alg` name, aligning it with `get_unverified_header`. No other + unverified-decode method enforces `alg` recognition, and unverified + decode was never a security boundary. + ## 0.7.0 — 2026-08-26 Performance release. Verified `decode` is about **2.3× faster** and `encode` about diff --git a/pyproject.toml b/pyproject.toml index b8ab20d..d9d71cd 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -4,7 +4,7 @@ build-backend = "maturin" [project] name = "oxyjwt" -version = "0.7.0" +version = "0.8.0" description = "High-performance Python JWT/JWS library backed by a Rust core (PyJWT drop-in replacement)." readme = "README.md" requires-python = ">=3.10" diff --git a/python/oxyjwt/__init__.py b/python/oxyjwt/__init__.py index 93322f3..4f76876 100644 --- a/python/oxyjwt/__init__.py +++ b/python/oxyjwt/__init__.py @@ -1,6 +1,6 @@ """OxyJWT public API (PyJWT-shaped module surface).""" -__version__ = "0.7.0" +__version__ = "0.8.0" from ._oxyjwt import ( DecodingKey, diff --git a/python/oxyjwt/_oxyjwt.pyi b/python/oxyjwt/_oxyjwt.pyi index cb7323e..af9e74f 100644 --- a/python/oxyjwt/_oxyjwt.pyi +++ b/python/oxyjwt/_oxyjwt.pyi @@ -92,7 +92,7 @@ def decode_verified_complete( def jws_parse_compact( token: str, -) -> tuple[bytes, dict[str, Any], bytes, bytes]: ... +) -> tuple[dict[str, Any], bytes, bytes]: ... def get_unverified_header(token: str) -> dict[str, Any]: ... diff --git a/python/oxyjwt/api_jwt.py b/python/oxyjwt/api_jwt.py index 8381d46..ecfd86c 100644 --- a/python/oxyjwt/api_jwt.py +++ b/python/oxyjwt/api_jwt.py @@ -115,40 +115,6 @@ def _check_string_or_iterable(value: object, name: str) -> None: raise TypeError(f"{name} must be a string, iterable or None") -def _is_plain_decode( - options: dict[str, Any] | None, - verify: bool | None, - detached_payload: bytes | None, - audience: object, - issuer: object, - subject: str | None, - leeway: float | timedelta, - typ: str | None, - algorithms: list[str] | None, -) -> bool: - """True when no argument changes the default verified-decode behaviour. - - With nothing but ``algorithms`` supplied, Rust validates ``exp``/``nbf`` - and has no audience/issuer/subject to check, so the remaining claim rules - reduce to :meth:`PyJWT._validate_claims_default`. A ``timedelta`` leeway - never compares equal to ``0`` and therefore takes the general path, as do - algorithm containers the native module cannot read directly (a set or an - iterator), which the general path normalises with ``list()``. - """ - return ( - options is None - and audience is None - and issuer is None - and subject is None - and typ is None - and detached_payload is None - and verify is None - and leeway == 0 - and bool(algorithms) - and isinstance(algorithms, (list, tuple)) - ) - - def _json_default_from_encoder(encoder_cls: type[JSONEncoder]) -> Callable[[Any], Any]: enc = encoder_cls() @@ -222,18 +188,47 @@ def decode( typ: str | None = None, **kwargs: Any, ) -> Any: + # Inlined instead of a shared `_is_plain_decode` helper: this condition + # runs on every call, and a 9-argument function call costs more than + # the comparisons themselves. True when no argument changes the + # default verified-decode behaviour: with nothing but `algorithms` + # supplied, Rust validates exp/nbf and has no audience/issuer/subject + # to check, so the remaining claim rules reduce to + # `_validate_claims_default`. A `timedelta` leeway never compares + # equal to `0` and therefore takes the general path, as do algorithm + # containers the native module cannot read directly (a set or an + # iterator), which the general path normalises with `list()`. if ( self._default_options and not kwargs - and _is_plain_decode( - options, verify, detached_payload, audience, issuer, subject, leeway, - typ, algorithms, - ) + and options is None + and audience is None + and issuer is None + and subject is None + and typ is None + and detached_payload is None + and verify is None + and leeway == 0 + and bool(algorithms) + and isinstance(algorithms, (list, tuple)) ): token = jwt if isinstance(jwt, str) else jwt.decode("utf-8") - _require_detached_payload_for_rfc7797( - token, detached_payload=None, verify_signature=True - ) + # Inlined `_require_detached_payload_for_rfc7797` for the same + # reason: detached_payload/verify_signature are always None/True + # here, so its own early-return check is dead code on this path. + # ".." can only appear via an empty (b64:false) payload segment, + # since base64url never emits a dot; that keeps ordinary tokens + # off the split/header-parse path below. + if ".." in token: + segments = token.split(".", 2) + if len(segments) >= 3 and segments[1] == "": + header_peek = _as_plain_dict(_oxyjwt.get_unverified_header(token)) + if header_peek.get("b64") is False: + raise DecodeError( + 'It is required that you pass in a value for the ' + '"detached_payload" argument to decode a message ' + "having the b64 header set to false." + ) payload = _oxyjwt.decode(token, key, algorithms) self._validate_claims_default(payload) return payload @@ -275,18 +270,33 @@ def decode_complete( typ: str | None = None, **kwargs: Any, ) -> dict[str, Any]: + # See the matching block in `decode()` for why this is inlined rather + # than calling a shared `_is_plain_decode` / rfc7797-peek helper. if ( self._default_options and not kwargs - and _is_plain_decode( - options, verify, detached_payload, audience, issuer, subject, leeway, - typ, algorithms, - ) + and options is None + and audience is None + and issuer is None + and subject is None + and typ is None + and detached_payload is None + and verify is None + and leeway == 0 + and bool(algorithms) + and isinstance(algorithms, (list, tuple)) ): token = jwt if isinstance(jwt, str) else jwt.decode("utf-8") - _require_detached_payload_for_rfc7797( - token, detached_payload=None, verify_signature=True - ) + if ".." in token: + segments = token.split(".", 2) + if len(segments) >= 3 and segments[1] == "": + header_peek = _as_plain_dict(_oxyjwt.get_unverified_header(token)) + if header_peek.get("b64") is False: + raise DecodeError( + 'It is required that you pass in a value for the ' + '"detached_payload" argument to decode a message ' + "having the b64 header set to false." + ) payload, header, signature = _oxyjwt.decode_verified_complete( token, key, algorithms ) @@ -377,7 +387,7 @@ def decode_complete( lwf = _leeway_seconds(leeway) if not verify_signature: - _s, header_obj, pld_bytes, sigb = _oxyjwt.jws_parse_compact(token) + header_obj, pld_bytes, sigb = _oxyjwt.jws_parse_compact(token) header = _as_plain_dict(header_obj) if detached_payload is not None: pld_bytes = bytes(detached_payload) @@ -454,30 +464,42 @@ def _validate_claims_default(self, payload: dict[str, Any]) -> None: ``rust_standard_claims`` true: Rust already validated ``exp``/``nbf``, so only ``iat``, the Python ``exp`` boundary, the "no audience expected" rule and the ``sub`` type check remain. - """ - now = time.time() + ``exp`` is only re-checked here when it is not a plain ``int``. Rust + already enforced ``exp > now`` using an integer clock (leeway is 0 on + this path), and for an integer ``exp`` that is the exact same + predicate as the truncating check below, so redoing it would only + cost a ``time.time()`` call for no behavioural difference. A ``float`` + ``exp`` still needs it: Rust rounds a fractional value to the nearest + second, while PyJWT (and the check below, for parity) truncates it, + so the two can disagree right at the boundary. + """ iat = payload.get("iat", _MISSING) - if iat is not _MISSING: - try: - iat_value = int(iat) - except (ValueError, TypeError) as e: - raise InvalidIssuedAtError( - "Issued At claim (iat) must be an integer." - ) from e - if iat_value > now: - raise ImmatureSignatureError("The token is not yet valid (iat)") - exp = payload.get("exp", _MISSING) - if exp is not _MISSING: - try: - exp_value = int(exp) - except (ValueError, TypeError) as e: - raise DecodeError( - "Expiration Time claim (exp) must be an integer." - ) from e - if exp_value <= now: - raise ExpiredSignatureError("Signature has expired") + exp_needs_check = exp is not _MISSING and type(exp) is not int + + if iat is not _MISSING or exp_needs_check: + now = time.time() + + if iat is not _MISSING: + try: + iat_value = int(iat) + except (ValueError, TypeError) as e: + raise InvalidIssuedAtError( + "Issued At claim (iat) must be an integer." + ) from e + if iat_value > now: + raise ImmatureSignatureError("The token is not yet valid (iat)") + + if exp_needs_check: + try: + exp_value = int(exp) + except (ValueError, TypeError) as e: + raise DecodeError( + "Expiration Time claim (exp) must be an integer." + ) from e + if exp_value <= now: + raise ExpiredSignatureError("Signature has expired") if payload.get("aud"): raise InvalidAudienceError("Invalid audience") diff --git a/rust/Cargo.lock b/rust/Cargo.lock index aab2071..4bcb91a 100644 --- a/rust/Cargo.lock +++ b/rust/Cargo.lock @@ -496,8 +496,9 @@ checksum = "9f7c3e4beb33f85d45ae3e3a1792185706c8e16d043238c593331cc7cd313b50" [[package]] name = "oxyjwt" -version = "0.7.0" +version = "0.8.0" dependencies = [ + "aws-lc-rs", "base64", "jsonwebtoken", "pyo3", diff --git a/rust/Cargo.toml b/rust/Cargo.toml index 74c62f6..5717ed8 100644 --- a/rust/Cargo.toml +++ b/rust/Cargo.toml @@ -1,6 +1,6 @@ [package] name = "oxyjwt" -version = "0.7.0" +version = "0.8.0" edition = "2021" rust-version = "1.85" license = "MIT" @@ -13,7 +13,7 @@ crate-type = ["cdylib"] [features] default = ["aws_lc_rs"] rust_crypto = ["jsonwebtoken/rust_crypto"] -aws_lc_rs = ["jsonwebtoken/aws_lc_rs"] +aws_lc_rs = ["jsonwebtoken/aws_lc_rs", "dep:aws-lc-rs"] [dependencies] base64 = "0.22" @@ -22,6 +22,12 @@ pyo3 = { version = "0.28", features = ["abi3-py310"] } serde_json = "1" zeroize = "1" +# Used directly (not just through jsonwebtoken) so RSA/RSA-PSS private keys can be +# parsed once at `EncodingKey.from_rsa_pem` time and signed with repeatedly, instead +# of jsonwebtoken re-parsing (and re-validating) the DER key on every `encode` call. +# See rust/src/keys.rs and the `cached_rsa_*` helpers in rust/src/api.rs. +aws-lc-rs = { version = "1", optional = true } + # Release wheels are built once and then run in a hot path, so trade build time # for codegen quality. `panic = "abort"` is deliberately not set: PyO3 unwinds # across the extension boundary to turn Rust panics into Python exceptions. diff --git a/rust/src/api.rs b/rust/src/api.rs index 34f8492..04a620a 100644 --- a/rust/src/api.rs +++ b/rust/src/api.rs @@ -2,14 +2,15 @@ use base64::engine::general_purpose::URL_SAFE_NO_PAD; use base64::Engine; use jsonwebtoken::errors::{new_error, ErrorKind}; use jsonwebtoken::{ - crypto::verify as jwt_crypto_verify, dangerous, encode as jwt_encode, DecodingKey, Header, + crypto::verify as jwt_crypto_verify, encode as jwt_encode, DecodingKey, Header, }; use pyo3::prelude::*; use pyo3::types::PyBytes; use serde_json::Value; use crate::algorithms::{ - algorithm_name, ensure_single_family, parse_algorithm, parse_algorithm_name, + algorithm_family, algorithm_name, ensure_single_family, parse_algorithm, parse_algorithm_name, + KeyFamily, }; use crate::claims::{json_to_py, py_to_json_for_encode}; use crate::claims_validate; @@ -18,12 +19,15 @@ use crate::jws; use crate::keys::{decoding_key_from_py, encoding_key_from_py}; use crate::validation::{self, DecodeValidation}; -/// Size limit plus strict three-segment compact JWS check (before parsing). -fn ensure_valid_compact_jwt(token: &str) -> PyResult<()> { - jws::split_compact_segments(token) - .map(|_| ()) - .map_err(errors::decode_error) -} +#[cfg(feature = "aws_lc_rs")] +use crate::keys::cached_rsa_encoding_key_from_py; +#[cfg(feature = "aws_lc_rs")] +use aws_lc_rs::rand::SystemRandom; +#[cfg(feature = "aws_lc_rs")] +use aws_lc_rs::signature::{ + RsaEncoding, RsaKeyPair, RSA_PKCS1_SHA256, RSA_PKCS1_SHA384, RSA_PKCS1_SHA512, RSA_PSS_SHA256, + RSA_PSS_SHA384, RSA_PSS_SHA512, +}; /// Failure from a decode step that runs with the GIL released. /// @@ -59,7 +63,15 @@ struct VerifiedToken { claims: Value, } -/// Verify a compact JWT and parse it in a single pass. +/// Verify a compact JWT and parse it in a single pass, optionally also +/// returning the decoded signature bytes for callers that need them +/// (`decode_complete`). `jwt_crypto_verify` already base64-decodes the +/// signature segment internally to check it; when `with_signature` is set, +/// this decodes it a second time to hand the bytes back, rather than making +/// the caller re-split the whole token and decode the signature segment +/// itself afterwards (as a separate `extract_signature_bytes` call used to). +/// A second small, bounded decode of the signature segment alone is the +/// trade-off for not re-scanning the (potentially much larger) full token. /// /// `jsonwebtoken::decode` parses the header twice and the payload twice (once /// for the caller's type, once for its internal validation struct) and returns @@ -67,11 +79,12 @@ struct VerifiedToken { /// Doing the steps here keeps it to one header parse, one payload parse and one /// signature check, and lets us reuse the already-parsed header as the Python /// header dict. -fn verify_and_parse( +fn verify_and_parse_impl( token: &str, decoding_key: &DecodingKey, decode_validation: &DecodeValidation, -) -> Result { + with_signature: bool, +) -> Result<(VerifiedToken, Option>), DecodeFail> { let (header_segment, payload_segment, signature_segment) = jws::split_compact_segments(token).map_err(DecodeFail::Decode)?; @@ -89,6 +102,11 @@ fn verify_and_parse( return Err(decode_fail(ErrorKind::InvalidSignature)); } + let signature = with_signature + .then(|| URL_SAFE_NO_PAD.decode(signature_segment)) + .transpose() + .map_err(|err| DecodeFail::Decode(err.to_string()))?; + let payload = URL_SAFE_NO_PAD .decode(payload_segment) .map_err(|err| DecodeFail::Decode(err.to_string()))?; @@ -102,7 +120,29 @@ fn verify_and_parse( claims_validate::validate_claims_value(&claims, &decode_validation.validation)?; - Ok(VerifiedToken { header, claims }) + Ok((VerifiedToken { header, claims }, signature)) +} + +fn verify_and_parse( + token: &str, + decoding_key: &DecodingKey, + decode_validation: &DecodeValidation, +) -> Result { + verify_and_parse_impl(token, decoding_key, decode_validation, false) + .map(|(verified, _)| verified) +} + +fn verify_and_parse_with_signature( + token: &str, + decoding_key: &DecodingKey, + decode_validation: &DecodeValidation, +) -> Result<(VerifiedToken, Vec), DecodeFail> { + let (verified, signature) = + verify_and_parse_impl(token, decoding_key, decode_validation, true)?; + Ok(( + verified, + signature.expect("with_signature=true always returns Some"), + )) } fn header_algorithm(header: &Value) -> Result { @@ -123,6 +163,112 @@ fn ensure_single_family_for_decode( .map_err(|_| decode_fail(ErrorKind::InvalidAlgorithm)) } +/// True when every algorithm in the allow-list is HMAC. Checked against the +/// caller's allow-list rather than the algorithm actually used, because for +/// `decode` the latter is only known after parsing the token header, which +/// happens inside the (possibly GIL-attached) verification step itself; an +/// HMAC-only allow-list already guarantees the verified algorithm is HMAC +/// too (`ensure_single_family_for_decode` rejects a mixed-family list). +fn algorithms_are_all_hmac(algorithms: &[jsonwebtoken::Algorithm]) -> bool { + algorithms + .iter() + .all(|algorithm| algorithm_family(*algorithm) == KeyFamily::Hmac) +} + +/// Releases the GIL around `f` unless `skip_detach` is set, in which case `f` +/// runs while still attached. HMAC sign/verify is fast enough (roughly a +/// microsecond) that `py.detach`'s release-and-reacquire can cost as much as +/// the operation itself, for negligible concurrency benefit on that single +/// call; RSA/EC/EdDSA operations are one to several orders of magnitude +/// slower and keep releasing the GIL unconditionally. See #124. +fn maybe_detach(py: Python<'_>, skip_detach: bool, f: F) -> T +where + F: pyo3::marker::Ungil + FnOnce() -> T, + T: pyo3::marker::Ungil, +{ + if skip_detach { + f() + } else { + py.detach(f) + } +} + +/// Length of URL-safe, unpadded base64 output for `n` input bytes. +fn base64_len_no_pad(n: usize) -> usize { + (n / 3) * 4 + + match n % 3 { + 0 => 0, + 1 => 2, + _ => 3, + } +} + +/// Base64url-encode `header` and `payload` directly into one pre-sized +/// `header.payload` string, instead of encoding each half into its own +/// `String` (via `Engine::encode`) and then copying both of those into a +/// third one with `format!`. +fn signing_input_string(header: &[u8], payload: &[u8]) -> String { + let mut signing_input = String::with_capacity( + base64_len_no_pad(header.len()) + 1 + base64_len_no_pad(payload.len()), + ); + URL_SAFE_NO_PAD.encode_string(header, &mut signing_input); + signing_input.push('.'); + URL_SAFE_NO_PAD.encode_string(payload, &mut signing_input); + signing_input +} + +/// The `aws_lc_rs` padding/digest scheme for an RSA/RSA-PSS algorithm. +#[cfg(feature = "aws_lc_rs")] +fn rsa_padding_for(algorithm: jsonwebtoken::Algorithm) -> &'static dyn RsaEncoding { + use jsonwebtoken::Algorithm; + match algorithm { + Algorithm::RS256 => &RSA_PKCS1_SHA256, + Algorithm::RS384 => &RSA_PKCS1_SHA384, + Algorithm::RS512 => &RSA_PKCS1_SHA512, + Algorithm::PS256 => &RSA_PSS_SHA256, + Algorithm::PS384 => &RSA_PSS_SHA384, + Algorithm::PS512 => &RSA_PSS_SHA512, + other => unreachable!( + "cached RSA signer requested for non-RSA algorithm {other:?}; \ + cached_rsa_encoding_key_from_py only returns Some for RSA family keys" + ), + } +} + +/// Sign `message` with an already-parsed RSA key, skipping the per-call +/// `RsaKeyPair::from_der` parse (and key validation) that +/// `jsonwebtoken::crypto::sign` would otherwise redo on every call. See #120. +#[cfg(feature = "aws_lc_rs")] +fn sign_with_cached_rsa_key( + key_pair: &RsaKeyPair, + algorithm: jsonwebtoken::Algorithm, + message: &[u8], +) -> PyResult> { + let padding = rsa_padding_for(algorithm); + let mut signature = vec![0u8; key_pair.public_modulus_len()]; + let rng = SystemRandom::new(); + key_pair + .sign(padding, &rng, message, &mut signature) + .map_err(|_| errors::encode_error("failed to sign with RSA key"))?; + Ok(signature) +} + +/// Build a compact JWS (`header.payload.signature`) from already base64url-ready +/// JSON bytes, signing with a cached, pre-parsed RSA key. +#[cfg(feature = "aws_lc_rs")] +fn sign_compact_with_cached_rsa( + header_json: &[u8], + payload_json: &[u8], + algorithm: jsonwebtoken::Algorithm, + key_pair: &RsaKeyPair, +) -> PyResult { + let mut token = signing_input_string(header_json, payload_json); + let signature = sign_with_cached_rsa_key(key_pair, algorithm, token.as_bytes())?; + token.push('.'); + URL_SAFE_NO_PAD.encode_string(signature, &mut token); + Ok(token) +} + #[pyfunction] #[pyo3(signature = (payload, key, algorithm = "HS256", headers = None))] pub fn encode( @@ -140,10 +286,25 @@ pub fn encode( let mut header = Header::new(algorithm); apply_headers(&mut header, headers, algorithm)?; + + #[cfg(feature = "aws_lc_rs")] + if let Some(rsa_key) = cached_rsa_encoding_key_from_py(key, algorithm)? { + let header_json = serde_json::to_vec(&header) + .map_err(|e| errors::encode_error(format!("failed to serialize header: {e}")))?; + let claims_json = serde_json::to_vec(&claims) + .map_err(|e| errors::encode_error(format!("failed to serialize claims: {e}")))?; + return py.detach(move || { + sign_compact_with_cached_rsa(&header_json, &claims_json, algorithm, rsa_key) + }); + } + let encoding_key = encoding_key_from_py(key, algorithm)?; + let skip_detach = algorithm_family(algorithm) == KeyFamily::Hmac; - py.detach(|| jwt_encode(&header, &claims, &encoding_key)) - .map_err(errors::from_jwt_encode_error) + maybe_detach(py, skip_detach, || { + jwt_encode(&header, &claims, &encoding_key) + }) + .map_err(errors::from_jwt_encode_error) } #[pyfunction] @@ -176,10 +337,12 @@ pub fn decode( algorithms, audience, issuer, subject, leeway, options, require, )?; let decoding_key = decoding_key_from_py(key, decode_validation.algorithms())?; + let skip_detach = algorithms_are_all_hmac(decode_validation.algorithms()); - let verified = py - .detach(|| verify_and_parse(token, &decoding_key, &decode_validation)) - .map_err(map_decode_fail)?; + let verified = maybe_detach(py, skip_detach, || { + verify_and_parse(token, &decoding_key, &decode_validation) + }) + .map_err(map_decode_fail)?; json_to_py(py, &verified.claims) } @@ -231,13 +394,11 @@ pub fn decode_verified_complete( ); } - let (verified, signature) = py - .detach(|| -> Result<(VerifiedToken, Vec), DecodeFail> { - let verified = verify_and_parse(token, &decoding_key, &decode_validation)?; - let signature = jws::extract_signature_bytes(token).map_err(DecodeFail::Decode)?; - Ok((verified, signature)) - }) - .map_err(map_decode_fail)?; + let skip_detach = algorithms_are_all_hmac(decode_validation.algorithms()); + let (verified, signature) = maybe_detach(py, skip_detach, || { + verify_and_parse_with_signature(token, &decoding_key, &decode_validation) + }) + .map_err(map_decode_fail)?; let claims_py = json_to_py(py, &verified.claims)?; let header_py = json_to_py(py, &verified.header)?; @@ -259,41 +420,41 @@ fn decode_rfc7797_verified_complete( ))); } - let (parts, claims, signature) = py - .detach(|| -> Result<_, DecodeFail> { - let parts = jws::parse_rfc7797_compact(token).map_err(DecodeFail::Token)?; - let claims: Value = serde_json::from_slice(detached_payload) - .map_err(|err| DecodeFail::Decode(format!("Invalid payload string: {err}")))?; - if !claims.is_object() { - return Err(DecodeFail::Decode( - "Invalid payload string: must be a json object".to_owned(), - )); - } + let skip_detach = algorithms_are_all_hmac(decode_validation.algorithms()); + let (parts, claims, signature) = maybe_detach(py, skip_detach, || -> Result<_, DecodeFail> { + let parts = jws::parse_rfc7797_compact(token).map_err(DecodeFail::Token)?; + let claims: Value = serde_json::from_slice(detached_payload) + .map_err(|err| DecodeFail::Decode(format!("Invalid payload string: {err}")))?; + if !claims.is_object() { + return Err(DecodeFail::Decode( + "Invalid payload string: must be a json object".to_owned(), + )); + } - let algorithm = header_algorithm(&parts.header)?; - if !decode_validation.algorithms().contains(&algorithm) { - return Err(decode_fail(ErrorKind::InvalidAlgorithm)); - } - ensure_single_family_for_decode(decode_validation.algorithms())?; - - let signing_input = jws::signing_input_rfc7797(&parts.header_segment, detached_payload); - if !jwt_crypto_verify( - &parts.signature_segment, - &signing_input, - decoding_key, - algorithm, - )? { - return Err(decode_fail(ErrorKind::InvalidSignature)); - } + let algorithm = header_algorithm(&parts.header)?; + if !decode_validation.algorithms().contains(&algorithm) { + return Err(decode_fail(ErrorKind::InvalidAlgorithm)); + } + ensure_single_family_for_decode(decode_validation.algorithms())?; + + let signing_input = jws::signing_input_rfc7797(&parts.header_segment, detached_payload); + if !jwt_crypto_verify( + &parts.signature_segment, + &signing_input, + decoding_key, + algorithm, + )? { + return Err(decode_fail(ErrorKind::InvalidSignature)); + } - claims_validate::validate_claims_value(&claims, &decode_validation.validation)?; + claims_validate::validate_claims_value(&claims, &decode_validation.validation)?; - let signature = URL_SAFE_NO_PAD - .decode(&parts.signature_segment) - .map_err(|err| DecodeFail::Decode(err.to_string()))?; - Ok((parts, claims, signature)) - }) - .map_err(map_decode_fail)?; + let signature = URL_SAFE_NO_PAD + .decode(&parts.signature_segment) + .map_err(|err| DecodeFail::Decode(err.to_string()))?; + Ok((parts, claims, signature)) + }) + .map_err(map_decode_fail)?; let header_py = json_to_py(py, &parts.header)?; let claims_py = json_to_py(py, &claims)?; @@ -303,10 +464,13 @@ fn decode_rfc7797_verified_complete( #[pyfunction] pub fn get_unverified_header(py: Python<'_>, token: &str) -> PyResult> { - ensure_valid_compact_jwt(token)?; - let token = token.to_owned(); + // `parse_compact_header_json` already runs `split_compact_segments` (size + // cap + strict three-segment check); a separate `ensure_valid_compact_jwt` + // pre-check would just repeat that scan of the whole token. `token: &str` + // needs no `to_owned()` either: `py.detach` only requires the closure to + // be `Ungil` (`Send`), which a `&str` already is, not `'static`. let header = py - .detach(move || jws::parse_compact_header_json(&token)) + .detach(move || jws::parse_compact_header_json(token)) .map_err(errors::decode_error)?; json_to_py(py, &header) @@ -314,13 +478,20 @@ pub fn get_unverified_header(py: Python<'_>, token: &str) -> PyResult> #[pyfunction] pub fn decode_unverified(py: Python<'_>, token: &str) -> PyResult> { - ensure_valid_compact_jwt(token)?; - let token = token.to_owned(); - let token_data = py - .detach(move || dangerous::insecure_decode::(&token)) - .map_err(errors::from_jwt_decode_error)?; + // `jsonwebtoken::dangerous::insecure_decode` re-implements its own + // lenient segment split (silently misparsing a token with extra `.`s + // rather than rejecting it) and, being generic over the return type, + // fully deserializes the header into `jsonwebtoken`'s own `Header` + // struct even though only `.claims` is ever read here. Reusing our own + // strict, single-pass `jws` helpers instead means one split (with our + // size cap and segment-count check), a JSON-object check on the header + // to reject a malformed one, and a payload parse -- with the header + // value itself never even converted to a Python object. + let claims = py + .detach(move || jws::parse_compact_claims_unverified(token)) + .map_err(errors::decode_error)?; - json_to_py(py, &token_data.claims) + json_to_py(py, &claims) } fn apply_headers( @@ -434,25 +605,28 @@ pub fn encode_json( let algorithm = parse_algorithm(algorithm)?; let mut header = Header::new(algorithm); apply_headers(&mut header, headers, algorithm)?; - let encoding_key = encoding_key_from_py(key, algorithm)?; let header_json = serde_json::to_vec(&header) .map_err(|e| errors::encode_error(format!("failed to serialize header: {e}")))?; let payload_owned = payload_bytes.to_vec(); - py.detach(move || { - use base64::engine::general_purpose::URL_SAFE_NO_PAD; - use base64::Engine; - - let header_b64 = URL_SAFE_NO_PAD.encode(&header_json); - let payload_b64 = URL_SAFE_NO_PAD.encode(&payload_owned); - let signing_input = format!("{header_b64}.{payload_b64}"); - - let signature = - jsonwebtoken::crypto::sign(signing_input.as_bytes(), &encoding_key, algorithm) - .map_err(errors::from_jwt_encode_error)?; + #[cfg(feature = "aws_lc_rs")] + if let Some(rsa_key) = cached_rsa_encoding_key_from_py(key, algorithm)? { + return py.detach(move || { + sign_compact_with_cached_rsa(&header_json, &payload_owned, algorithm, rsa_key) + }); + } - Ok(format!("{signing_input}.{signature}")) + let encoding_key = encoding_key_from_py(key, algorithm)?; + let skip_detach = algorithm_family(algorithm) == KeyFamily::Hmac; + + maybe_detach(py, skip_detach, move || { + let mut token = signing_input_string(&header_json, &payload_owned); + let signature = jsonwebtoken::crypto::sign(token.as_bytes(), &encoding_key, algorithm) + .map_err(errors::from_jwt_encode_error)?; + token.push('.'); + token.push_str(&signature); + Ok(token) }) } diff --git a/rust/src/claims.rs b/rust/src/claims.rs index 0f56b58..266ee1b 100644 --- a/rust/src/claims.rs +++ b/rust/src/claims.rs @@ -95,6 +95,30 @@ fn py_to_json_depth(value: &Bound<'_, PyAny>, depth: usize) -> PyResult { )) } +/// Standard JWT claim / header names, interned once per key rather than +/// allocated fresh on every `set_item`. `pyo3::intern!` caches each key in a +/// call-site-local static, both skipping the per-call `PyString` allocation +/// and interning the string in CPython's own intern table, so a subsequent +/// Python-side `dict.get("exp")` (a `str` literal, which CPython also +/// interns) can hit the identity-comparison fast path during dict lookup +/// instead of a full string comparison. Anything outside this list falls +/// back to an ordinary, uninterned `PyString`. +fn claim_key<'py>(py: Python<'py>, key: &str) -> Bound<'py, PyString> { + match key { + "exp" => pyo3::intern!(py, "exp").clone(), + "iat" => pyo3::intern!(py, "iat").clone(), + "nbf" => pyo3::intern!(py, "nbf").clone(), + "sub" => pyo3::intern!(py, "sub").clone(), + "aud" => pyo3::intern!(py, "aud").clone(), + "iss" => pyo3::intern!(py, "iss").clone(), + "jti" => pyo3::intern!(py, "jti").clone(), + "alg" => pyo3::intern!(py, "alg").clone(), + "typ" => pyo3::intern!(py, "typ").clone(), + "kid" => pyo3::intern!(py, "kid").clone(), + _ => PyString::new(py, key), + } +} + fn json_to_bound<'py>(py: Python<'py>, value: &Value) -> PyResult> { match value { Value::Null => Ok(py.None().into_bound(py)), @@ -122,7 +146,7 @@ fn json_to_bound<'py>(py: Python<'py>, value: &Value) -> PyResult { let dict = PyDict::new(py); for (key, value) in values { - dict.set_item(key, json_to_bound(py, value)?)?; + dict.set_item(claim_key(py, key), json_to_bound(py, value)?)?; } Ok(dict.into_any()) } diff --git a/rust/src/jws.rs b/rust/src/jws.rs index 91838a9..0dea053 100644 --- a/rust/src/jws.rs +++ b/rust/src/jws.rs @@ -8,7 +8,7 @@ use serde_json::Value; use crate::errors; -type CompactJwsParts = (Vec, Value, Vec, Vec); +type CompactJwsParts = (Value, Vec, Vec); /// Maximum compact serialization size (`header.payload.signature`) before parsing. pub const MAX_COMPACT_JWT_BYTES: usize = 256 * 1024; @@ -136,39 +136,48 @@ pub fn signing_input_rfc7797(header_segment: &str, payload: &[u8]) -> Vec { signing_input } -/// Returns `(signing_input bytes, header JSON object, raw payload bytes, signature bytes)`. +/// Returns `(header JSON object, raw payload bytes, signature bytes)`. +/// +/// Does not compute the `header.payload` signing input: `jws_parse_compact`'s +/// only Python caller (the `decode_complete(verify_signature=False)` path in +/// `api_jwt.py`) never used it, so returning it was a wasted `Vec` clone +/// of the token's own bytes on every unverified `decode_complete` call. pub fn parse_compact_jws(token: &str) -> Result { let (h, p, s) = split_compact_segments(token)?; - let signing_input_len = h.len().saturating_add(1).saturating_add(p.len()); - let signing_input = token.as_bytes()[..signing_input_len].to_vec(); let header = decode_header_json(h)?; let payload_bytes = URL_SAFE_NO_PAD.decode(p).map_err(|e| e.to_string())?; let signature_bytes = URL_SAFE_NO_PAD.decode(s).map_err(|e| e.to_string())?; - Ok((signing_input, header, payload_bytes, signature_bytes)) + Ok((header, payload_bytes, signature_bytes)) } -/// Extract and decode the JWS signature segment without parsing header or payload JSON. -pub fn extract_signature_bytes(token: &str) -> Result, String> { - let (_, _, sig_encoded) = split_compact_segments(token)?; - URL_SAFE_NO_PAD - .decode(sig_encoded) - .map_err(|e| e.to_string()) +/// Decode only the claims (payload) segment of a compact JWT, for the +/// module-level `decode_unverified`, which never looks at the header or +/// signature: no reason to spend a JSON parse (or a base64 decode, for the +/// signature) on either. Still runs `split_compact_segments`, so the same +/// size cap and strict three-segment check apply as everywhere else; the +/// header is parsed as JSON only to reject a malformed one, matching +/// `get_unverified_header`'s validation, then discarded. +pub fn parse_compact_claims_unverified(token: &str) -> Result { + let (header_segment, payload_segment, _) = split_compact_segments(token)?; + decode_header_json(header_segment)?; + let payload = URL_SAFE_NO_PAD + .decode(payload_segment) + .map_err(|e| e.to_string())?; + serde_json::from_slice(&payload).map_err(|e| e.to_string()) } -type JwsParseOutput = (Py, Py, Py, Py); +type JwsParseOutput = (Py, Py, Py); #[pyfunction] pub fn jws_parse_compact(py: Python<'_>, token: &str) -> PyResult { - let token = token.to_owned(); - let (signing_input, header, payload, signature) = py - .detach(move || parse_compact_jws(&token)) + let (header, payload, signature) = py + .detach(move || parse_compact_jws(token)) .map_err(errors::decode_error)?; use crate::claims::json_to_py; let header_obj = json_to_py(py, &header)?; - let signing = PyBytes::new(py, &signing_input); let pld = PyBytes::new(py, &payload); let sigb = PyBytes::new(py, &signature); - Ok((signing.into(), header_obj, pld.into(), sigb.into())) + Ok((header_obj, pld.into(), sigb.into())) } #[cfg(test)] @@ -201,23 +210,32 @@ mod tests { "Too many segments" ); assert_eq!(parse_compact_jws(token).unwrap_err(), "Too many segments"); + assert_eq!( + parse_compact_claims_unverified(token).unwrap_err(), + "Too many segments" + ); } #[test] - fn borrowed_signing_input_matches_owned_parse() { + fn signing_input_of_matches_token_prefix() { let token = "eyJhbGciOiJIUzI1NiJ9.eyJzdWIiOiJ1In0.dozjgNryP4J3jVmNHl0w5N_XgL0n3I9PlFUP0THsR8U"; let (h, p, _) = split_compact_segments(token).expect("split"); - let (owned, _, _, _) = parse_compact_jws(token).expect("parse"); - assert_eq!(signing_input_of(token, h, p), owned.as_slice()); + assert_eq!( + signing_input_of(token, h, p), + b"eyJhbGciOiJIUzI1NiJ9.eyJzdWIiOiJ1In0".as_slice() + ); } #[test] - fn extract_signature_matches_full_parse() { + fn claims_unverified_matches_full_parse_payload() { let token = "eyJhbGciOiJIUzI1NiJ9.eyJzdWIiOiJ1In0.dozjgNryP4J3jVmNHl0w5N_XgL0n3I9PlFUP0THsR8U"; - let (_, _, _, full_sig) = parse_compact_jws(token).expect("parse"); - let extracted = extract_signature_bytes(token).expect("extract"); - assert_eq!(full_sig, extracted); + let (_, payload, _) = parse_compact_jws(token).expect("parse"); + let claims: Value = serde_json::from_slice(&payload).expect("payload is json"); + assert_eq!( + claims, + parse_compact_claims_unverified(token).expect("claims") + ); } } diff --git a/rust/src/keys.rs b/rust/src/keys.rs index 8ac808a..90b0ebe 100644 --- a/rust/src/keys.rs +++ b/rust/src/keys.rs @@ -7,10 +7,21 @@ use crate::algorithms::{ensure_algorithm_family, ensure_single_family, KeyFamily use crate::claims; use crate::errors; +#[cfg(feature = "aws_lc_rs")] +use aws_lc_rs::signature::RsaKeyPair; + #[derive(Debug)] struct EncodingKeyMaterial { family: KeyFamily, key: JwtEncodingKey, + // Populated only for RSA/RSA-PSS keys when the `aws_lc_rs` crypto backend is in + // use. `jsonwebtoken::crypto::sign` re-parses (and re-validates, which for RSA + // is the expensive part) `key.inner()` into an `aws_lc_rs::signature::RsaKeyPair` + // on every call; parsing it once here and signing through it directly in + // `api::encode` / `api::encode_json` turns that per-call cost into a one-time + // cost at key construction. See issue #120. + #[cfg(feature = "aws_lc_rs")] + cached_rsa: Option, } #[derive(Debug)] @@ -19,12 +30,19 @@ struct DecodingKeyMaterial { key: JwtDecodingKey, } -#[pyclass(module = "oxyjwt._oxyjwt")] +// `frozen` means Python cannot mutate the object after construction, so pyo3 +// hands out `&EncodingKey` / `&DecodingKey` via `Bound::get()` with no runtime +// borrow-flag check, and that reference's lifetime is tied to the calling +// scope rather than to a `PyRef` guard. `encoding_key_from_py` / +// `decoding_key_from_py` use this to borrow the underlying key material +// straight through to `py.detach` instead of cloning it on every +// `encode`/`decode` call. See issue #121. +#[pyclass(module = "oxyjwt._oxyjwt", frozen)] pub struct EncodingKey { material: EncodingKeyMaterial, } -#[pyclass(module = "oxyjwt._oxyjwt")] +#[pyclass(module = "oxyjwt._oxyjwt", frozen)] pub struct DecodingKey { material: DecodingKeyMaterial, } @@ -45,12 +63,15 @@ impl EncodingKey { #[staticmethod] pub fn from_rsa_pem(pem: &Bound<'_, PyAny>) -> PyResult { let bytes = bytes_from_py(pem)?; - Ok(Self { - material: EncodingKeyMaterial::new( - KeyFamily::Rsa, - JwtEncodingKey::from_rsa_pem(&bytes).map_err(errors::from_jwt_encode_error)?, - ), - }) + let key = JwtEncodingKey::from_rsa_pem(&bytes).map_err(errors::from_jwt_encode_error)?; + #[cfg(feature = "aws_lc_rs")] + let material = { + let cached_rsa = parse_cached_rsa_key_pair(key.inner())?; + EncodingKeyMaterial::new_rsa(key, cached_rsa) + }; + #[cfg(not(feature = "aws_lc_rs"))] + let material = EncodingKeyMaterial::new(KeyFamily::Rsa, key); + Ok(Self { material }) } #[staticmethod] @@ -154,12 +175,35 @@ impl DecodingKey { impl EncodingKeyMaterial { fn new(family: KeyFamily, key: JwtEncodingKey) -> Self { - Self { family, key } + Self { + family, + key, + #[cfg(feature = "aws_lc_rs")] + cached_rsa: None, + } + } + + #[cfg(feature = "aws_lc_rs")] + fn new_rsa(key: JwtEncodingKey, cached_rsa: RsaKeyPair) -> Self { + Self { + family: KeyFamily::Rsa, + key, + cached_rsa: Some(cached_rsa), + } + } + + fn encoding_key_ref(&self, algorithm: Algorithm) -> PyResult<&JwtEncodingKey> { + ensure_algorithm_family(algorithm, self.family)?; + Ok(&self.key) } - fn encoding_key(&self, algorithm: Algorithm) -> PyResult { + /// The pre-parsed RSA signing key, when `key` is an RSA/RSA-PSS `EncodingKey` + /// and `algorithm` is compatible with it. `None` for every other key family, + /// in which case the caller falls back to `encoding_key_ref` + `jsonwebtoken::crypto`. + #[cfg(feature = "aws_lc_rs")] + fn cached_rsa_signer_ref(&self, algorithm: Algorithm) -> PyResult> { ensure_algorithm_family(algorithm, self.family)?; - Ok(self.key.clone()) + Ok(self.cached_rsa.as_ref()) } } @@ -181,18 +225,57 @@ impl DecodingKeyMaterial { Ok(()) } - fn decoding_key(&self, algorithms: &[Algorithm]) -> PyResult { + fn decoding_key_ref(&self, algorithms: &[Algorithm]) -> PyResult<&JwtDecodingKey> { self.validate_for_algorithms(algorithms)?; - Ok(self.key.clone()) + Ok(&self.key) } } -pub fn encoding_key_from_py( - key: &Bound<'_, PyAny>, +/// Either a `JwtEncodingKey` borrowed straight out of a frozen `EncodingKey` +/// pyclass (the common case: a pre-built typed key reused across calls), or +/// one built on the spot from a raw HMAC secret. `Deref`s to `JwtEncodingKey` +/// so call sites use it exactly like an owned key. +pub enum BorrowedEncodingKey<'a> { + Ref(&'a JwtEncodingKey), + Owned(JwtEncodingKey), +} + +impl std::ops::Deref for BorrowedEncodingKey<'_> { + type Target = JwtEncodingKey; + + fn deref(&self) -> &JwtEncodingKey { + match self { + Self::Ref(key) => key, + Self::Owned(key) => key, + } + } +} + +/// Same as [`BorrowedEncodingKey`] for `JwtDecodingKey`. +pub enum BorrowedDecodingKey<'a> { + Ref(&'a JwtDecodingKey), + Owned(JwtDecodingKey), +} + +impl std::ops::Deref for BorrowedDecodingKey<'_> { + type Target = JwtDecodingKey; + + fn deref(&self) -> &JwtDecodingKey { + match self { + Self::Ref(key) => key, + Self::Owned(key) => key, + } + } +} + +pub fn encoding_key_from_py<'a>( + key: &'a Bound<'_, PyAny>, algorithm: Algorithm, -) -> PyResult { - if let Ok(key_ref) = key.extract::>() { - return key_ref.material.encoding_key(algorithm); +) -> PyResult> { + if let Ok(bound) = key.cast::() { + return Ok(BorrowedEncodingKey::Ref( + bound.get().material.encoding_key_ref(algorithm)?, + )); } if crate::algorithms::algorithm_family(algorithm) != KeyFamily::Hmac { @@ -202,24 +285,56 @@ pub fn encoding_key_from_py( } let bytes = secret_bytes_from_py(key)?; - Ok(JwtEncodingKey::from_secret(bytes.as_ref())) + Ok(BorrowedEncodingKey::Owned(JwtEncodingKey::from_secret( + bytes.as_ref(), + ))) } -pub fn decoding_key_from_py( - key: &Bound<'_, PyAny>, +/// Parse a DER-encoded RSA private key into an `aws_lc_rs` `RsaKeyPair`, mapping +/// failures to the same `InvalidKeyError` that `jsonwebtoken`'s own RSA key +/// rejection produces. +#[cfg(feature = "aws_lc_rs")] +fn parse_cached_rsa_key_pair(der: &[u8]) -> PyResult { + RsaKeyPair::from_der(der).map_err(|err| { + errors::from_jwt_encode_error(jsonwebtoken::errors::new_error( + jsonwebtoken::errors::ErrorKind::InvalidRsaKey(err.to_string()), + )) + }) +} + +/// The pre-parsed RSA signing key backing `key`, when `key` is an `EncodingKey` +/// built from `EncodingKey.from_rsa_pem` and `algorithm` is an RSA/RSA-PSS +/// algorithm compatible with it. `Ok(None)` when `key` is not an `EncodingKey` +/// object (a raw HMAC secret): the caller falls back to `encoding_key_from_py`, +/// which raises the appropriate error for that case. +#[cfg(feature = "aws_lc_rs")] +pub fn cached_rsa_encoding_key_from_py<'a>( + key: &'a Bound<'_, PyAny>, + algorithm: Algorithm, +) -> PyResult> { + match key.cast::() { + Ok(bound) => bound.get().material.cached_rsa_signer_ref(algorithm), + Err(_) => Ok(None), + } +} + +pub fn decoding_key_from_py<'a>( + key: &'a Bound<'_, PyAny>, algorithms: &[Algorithm], -) -> PyResult { - if let Ok(key_ref) = key.extract::>() { - return key_ref.material.decoding_key(algorithms); +) -> PyResult> { + if let Ok(bound) = key.cast::() { + return Ok(BorrowedDecodingKey::Ref( + bound.get().material.decoding_key_ref(algorithms)?, + )); } raw_decoding_key_from_py(key, algorithms) } -fn raw_decoding_key_from_py( +fn raw_decoding_key_from_py<'a>( key: &Bound<'_, PyAny>, algorithms: &[Algorithm], -) -> PyResult { +) -> PyResult> { let family = ensure_single_family(algorithms)?; if family != KeyFamily::Hmac { return Err(errors::invalid_key( @@ -228,7 +343,9 @@ fn raw_decoding_key_from_py( } let bytes = secret_bytes_from_py(key)?; - Ok(JwtDecodingKey::from_secret(bytes.as_ref())) + Ok(BorrowedDecodingKey::Owned(JwtDecodingKey::from_secret( + bytes.as_ref(), + ))) } /// Copy HMAC secret material from Python; buffer is zeroized on drop. diff --git a/tests/test_algorithms.py b/tests/test_algorithms.py index ba6179e..ad9a873 100644 --- a/tests/test_algorithms.py +++ b/tests/test_algorithms.py @@ -49,3 +49,24 @@ def test_supported_algorithm_names_are_recognized( decoded = oxyjwt.decode(token, ed_pair.decoding_key, algorithms=[algorithm]) assert decoded["exp"] == payload["exp"] + + +@pytest.mark.parametrize("algorithm", ["RS256", "RS384", "RS512", "PS256", "PS384", "PS512"]) +def test_rsa_encoding_key_signs_correctly_across_repeated_calls( + algorithm: str, rsa_pair: object +) -> None: + """`EncodingKey.from_rsa_pem` parses the private key once; every `encode` + call re-signs through the cached parsed key rather than re-parsing the PEM + (see issue #120), so the same `EncodingKey` object must keep producing + signatures that verify correctly across many calls, not just the first. + """ + for i in range(20): + payload = {"sub": f"user-{i}", "exp": int(time.time()) + 60} + token = oxyjwt.encode(payload, rsa_pair.encoding_key, algorithm=algorithm) + decoded = oxyjwt.decode(token, rsa_pair.decoding_key, algorithms=[algorithm]) + assert decoded == payload + + +def test_rsa_encoding_key_rejects_malformed_pem_at_construction() -> None: + with pytest.raises(oxyjwt.InvalidKeyError): + oxyjwt.EncodingKey.from_rsa_pem(b"not a pem encoded key") diff --git a/tests/test_decode_fast_path.py b/tests/test_decode_fast_path.py index 8f47170..b454f4a 100644 --- a/tests/test_decode_fast_path.py +++ b/tests/test_decode_fast_path.py @@ -146,6 +146,25 @@ def test_timedelta_leeway_still_honoured() -> None: )["sub"] == "u" +def test_float_exp_boundary_still_checked_in_python() -> None: + """An integer `exp` skips the redundant Python recheck (Rust's own + boundary check is exact for it), but a `float` `exp` must not. + + Pin `exp` to `K + 0.6` for the current whole second `K`: Rust rounds that + up to `K + 1` and, checked within the same second `K`, considers it not + yet expired. PyJWT (and the Python recheck below, for parity) truncates + it back down to `K`, which is already <= "now", so the token must still + be rejected. Adding a fraction to `time.time()` instead (`now + 0.6`) + would not reproduce this: it lands in the future either way and both + sides agree the token is valid. + """ + current_second = int(time.time()) + exp = current_second + 0.6 + token = oxyjwt.encode({"sub": "u", "exp": exp}, SECRET, "HS256") + with pytest.raises(oxyjwt.ExpiredSignatureError): + oxyjwt.decode(token, SECRET, algorithms=["HS256"]) + + class TestRequiredClaims: """`require` follows PyJWT: a JSON null counts as an absent claim.""" diff --git a/tests/test_encode_decode.py b/tests/test_encode_decode.py index 64c58a8..d17543e 100644 --- a/tests/test_encode_decode.py +++ b/tests/test_encode_decode.py @@ -1,6 +1,7 @@ from __future__ import annotations import time +from concurrent.futures import ThreadPoolExecutor from datetime import datetime, timezone import orjson @@ -172,6 +173,37 @@ def test_decode_unverified_is_explicit() -> None: assert oxyjwt.decode_unverified(token)["sub"] == "user-123" +def _b64u(data: bytes) -> str: + import base64 + + return base64.urlsafe_b64encode(data).rstrip(b"=").decode() + + +def test_decode_unverified_rejects_non_object_header() -> None: + """The header must still be a JSON object, even though its value is + never read: `decode_unverified` never parses the header into + `jsonwebtoken`'s typed `Header` struct, but still checks its shape. + """ + header = _b64u(b"not-json-at-all") + payload = _b64u(orjson.dumps({"sub": "u"})) + token = f"{header}.{payload}.sig" + with pytest.raises(oxyjwt.DecodeError): + oxyjwt.decode_unverified(token) + + +def test_decode_unverified_accepts_unrecognized_alg() -> None: + """`decode_unverified` only requires the header to be a JSON object, + matching `get_unverified_header`, not a recognized `alg` name: nothing + here is verified, so there is no security reason to be stricter about + the header than the sibling method that returns it. + """ + header = _b64u(orjson.dumps({"alg": "made-up-alg", "typ": "JWT"})) + payload = _b64u(orjson.dumps({"sub": "u"})) + token = f"{header}.{payload}.sig" + assert oxyjwt.decode_unverified(token) == {"sub": "u"} + assert oxyjwt.get_unverified_header(token) == {"alg": "made-up-alg", "typ": "JWT"} + + def test_encode_deeply_nested_claims_rejected() -> None: nested: dict[str, object] = {"a": 1} current = nested @@ -183,3 +215,58 @@ def test_encode_deeply_nested_claims_rejected() -> None: oxyjwt._oxyjwt.encode(nested, "secret", "HS256") +@pytest.mark.parametrize("claim", ["exp", "iat", "nbf", "sub", "aud", "iss", "jti"]) +def test_standard_claim_key_is_interned(claim: str) -> None: + """A standard claim name comes back as the same interned `str` object + pyo3 caches for it, matching a same-spelling literal by identity (`is`) + rather than just equality. + """ + token = oxyjwt.encode({claim: "v"}, "secret", "HS256") + payload = oxyjwt.decode_unverified(token) + (key,) = payload.keys() + assert key is claim + + +@pytest.mark.parametrize("field", ["alg", "typ", "kid"]) +def test_standard_header_key_is_interned(field: str) -> None: + token = oxyjwt.encode({"sub": "u"}, "secret", "HS256", headers={"kid": "k1"}) + header = oxyjwt.get_unverified_header(token) + (key,) = (k for k in header if k == field) + assert key is field + + +def test_custom_claim_key_is_not_interned() -> None: + """A key outside the standard list must still decode correctly, as an + ordinary (uninterned) `str` rather than being pulled from the cache. + """ + custom_key = "a_custom_claim" + token = oxyjwt.encode({custom_key: "v"}, "secret", "HS256") + payload = oxyjwt.decode_unverified(token) + (key,) = payload.keys() + assert key == custom_key + assert key is not custom_key + + +def test_hmac_encode_decode_are_thread_safe_without_gil_release() -> None: + """HS256 `encode`/`decode` no longer release the GIL around the native + call (see #124). That is purely a scheduling change on the Rust side, but + it is worth a concurrency regression test in its own right: each thread + must still get back exactly the token/payload it asked for, with no + cross-talk between threads sharing the same secret. + """ + secret = "concurrent-hmac-secret-with-plenty-of-length" + + def roundtrip(i: int) -> bool: + payload = {"sub": f"user-{i}", "n": i} + for _ in range(200): + token = oxyjwt.encode(payload, secret, "HS256") + decoded = oxyjwt.decode(token, secret, algorithms=["HS256"]) + if decoded != payload: + return False + return True + + with ThreadPoolExecutor(max_workers=8) as pool: + results = list(pool.map(roundtrip, range(32))) + assert all(results) + +