Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
17 changes: 16 additions & 1 deletion CHANGELOG.md
Original file line number Diff line number Diff line change
Expand Up @@ -7,7 +7,22 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0

## [Unreleased]

(No changes yet.)
### 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)

## [0.7.0] — 2026-08-26

Expand Down
1 change: 1 addition & 0 deletions rust/Cargo.lock

Some generated files are not rendered by default. Learn more about how customized files appear on GitHub.

8 changes: 7 additions & 1 deletion rust/Cargo.toml
Original file line number Diff line number Diff line change
Expand Up @@ -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"
Expand All @@ -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.
Expand Down
85 changes: 84 additions & 1 deletion rust/src/api.rs
Original file line number Diff line number Diff line change
Expand Up @@ -18,6 +18,16 @@ use crate::jws;
use crate::keys::{decoding_key_from_py, encoding_key_from_py};
use crate::validation::{self, DecodeValidation};

#[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,
};

/// Size limit plus strict three-segment compact JWS check (before parsing).
fn ensure_valid_compact_jwt(token: &str) -> PyResult<()> {
jws::split_compact_segments(token)
Expand Down Expand Up @@ -123,6 +133,59 @@ fn ensure_single_family_for_decode(
.map_err(|_| decode_fail(ErrorKind::InvalidAlgorithm))
}

/// 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<Vec<u8>> {
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<String> {
let header_b64 = URL_SAFE_NO_PAD.encode(header_json);
let payload_b64 = URL_SAFE_NO_PAD.encode(payload_json);
let signing_input = format!("{header_b64}.{payload_b64}");
let signature = sign_with_cached_rsa_key(key_pair, algorithm, signing_input.as_bytes())?;
let signature_b64 = URL_SAFE_NO_PAD.encode(signature);
Ok(format!("{signing_input}.{signature_b64}"))
}

#[pyfunction]
#[pyo3(signature = (payload, key, algorithm = "HS256", headers = None))]
pub fn encode(
Expand All @@ -140,6 +203,18 @@ 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)?;

py.detach(|| jwt_encode(&header, &claims, &encoding_key))
Expand Down Expand Up @@ -434,13 +509,21 @@ 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();

#[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)
});
}

let encoding_key = encoding_key_from_py(key, algorithm)?;

py.detach(move || {
use base64::engine::general_purpose::URL_SAFE_NO_PAD;
use base64::Engine;
Expand Down
82 changes: 75 additions & 7 deletions rust/src/keys.rs
Original file line number Diff line number Diff line change
Expand Up @@ -7,10 +7,24 @@ use crate::algorithms::{ensure_algorithm_family, ensure_single_family, KeyFamily
use crate::claims;
use crate::errors;

#[cfg(feature = "aws_lc_rs")]
use std::sync::Arc;

#[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<Arc<RsaKeyPair>>,
}

#[derive(Debug)]
Expand Down Expand Up @@ -45,12 +59,15 @@ impl EncodingKey {
#[staticmethod]
pub fn from_rsa_pem(pem: &Bound<'_, PyAny>) -> PyResult<Self> {
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 = Arc::new(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]
Expand Down Expand Up @@ -154,13 +171,36 @@ 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: Arc<RsaKeyPair>) -> Self {
Self {
family: KeyFamily::Rsa,
key,
cached_rsa: Some(cached_rsa),
}
}

fn encoding_key(&self, algorithm: Algorithm) -> PyResult<JwtEncodingKey> {
ensure_algorithm_family(algorithm, self.family)?;
Ok(self.key.clone())
}

/// 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` + `jsonwebtoken::crypto`.
#[cfg(feature = "aws_lc_rs")]
fn cached_rsa_signer(&self, algorithm: Algorithm) -> PyResult<Option<Arc<RsaKeyPair>>> {
ensure_algorithm_family(algorithm, self.family)?;
Ok(self.cached_rsa.clone())
}
}

impl DecodingKeyMaterial {
Expand Down Expand Up @@ -205,6 +245,34 @@ pub fn encoding_key_from_py(
Ok(JwtEncodingKey::from_secret(bytes.as_ref()))
}

/// 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> {
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(
key: &Bound<'_, PyAny>,
algorithm: Algorithm,
) -> PyResult<Option<Arc<RsaKeyPair>>> {
match key.extract::<PyRef<'_, EncodingKey>>() {
Ok(key_ref) => key_ref.material.cached_rsa_signer(algorithm),
Err(_) => Ok(None),
}
}

pub fn decoding_key_from_py(
key: &Bound<'_, PyAny>,
algorithms: &[Algorithm],
Expand Down
21 changes: 21 additions & 0 deletions tests/test_algorithms.py
Original file line number Diff line number Diff line change
Expand Up @@ -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")
Loading