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
8 changes: 8 additions & 0 deletions CHANGELOG.md
Original file line number Diff line number Diff line change
Expand Up @@ -23,6 +23,14 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0
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)

## [0.7.0] — 2026-08-26

Expand Down
4 changes: 2 additions & 2 deletions rust/src/api.rs
Original file line number Diff line number Diff line change
Expand Up @@ -211,7 +211,7 @@ pub fn encode(
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)
sign_compact_with_cached_rsa(&header_json, &claims_json, algorithm, rsa_key)
});
}

Expand Down Expand Up @@ -518,7 +518,7 @@ pub fn encode_json(
#[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)
sign_compact_with_cached_rsa(&header_json, &payload_owned, algorithm, rsa_key)
});
}

Expand Down
117 changes: 83 additions & 34 deletions rust/src/keys.rs
Original file line number Diff line number Diff line change
Expand Up @@ -7,9 +7,6 @@ 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;

Expand All @@ -24,7 +21,7 @@ struct EncodingKeyMaterial {
// `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>>,
cached_rsa: Option<RsaKeyPair>,
}

#[derive(Debug)]
Expand All @@ -33,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,
}
Expand All @@ -62,7 +66,7 @@ impl EncodingKey {
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())?);
let cached_rsa = parse_cached_rsa_key_pair(key.inner())?;
EncodingKeyMaterial::new_rsa(key, cached_rsa)
};
#[cfg(not(feature = "aws_lc_rs"))]
Expand Down Expand Up @@ -180,26 +184,26 @@ impl EncodingKeyMaterial {
}

#[cfg(feature = "aws_lc_rs")]
fn new_rsa(key: JwtEncodingKey, cached_rsa: Arc<RsaKeyPair>) -> Self {
fn new_rsa(key: JwtEncodingKey, cached_rsa: RsaKeyPair) -> Self {
Self {
family: KeyFamily::Rsa,
key,
cached_rsa: Some(cached_rsa),
}
}

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

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

Expand All @@ -221,18 +225,57 @@ impl DecodingKeyMaterial {
Ok(())
}

fn decoding_key(&self, algorithms: &[Algorithm]) -> PyResult<JwtDecodingKey> {
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<JwtEncodingKey> {
if let Ok(key_ref) = key.extract::<PyRef<'_, EncodingKey>>() {
return key_ref.material.encoding_key(algorithm);
) -> PyResult<BorrowedEncodingKey<'a>> {
if let Ok(bound) = key.cast::<EncodingKey>() {
return Ok(BorrowedEncodingKey::Ref(
bound.get().material.encoding_key_ref(algorithm)?,
));
}

if crate::algorithms::algorithm_family(algorithm) != KeyFamily::Hmac {
Expand All @@ -242,7 +285,9 @@ 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(),
)))
}

/// Parse a DER-encoded RSA private key into an `aws_lc_rs` `RsaKeyPair`, mapping
Expand All @@ -263,31 +308,33 @@ fn parse_cached_rsa_key_pair(der: &[u8]) -> PyResult<RsaKeyPair> {
/// 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>,
pub fn cached_rsa_encoding_key_from_py<'a>(
key: &'a Bound<'_, PyAny>,
algorithm: Algorithm,
) -> PyResult<Option<Arc<RsaKeyPair>>> {
match key.extract::<PyRef<'_, EncodingKey>>() {
Ok(key_ref) => key_ref.material.cached_rsa_signer(algorithm),
) -> PyResult<Option<&'a RsaKeyPair>> {
match key.cast::<EncodingKey>() {
Ok(bound) => bound.get().material.cached_rsa_signer_ref(algorithm),
Err(_) => Ok(None),
}
}

pub fn decoding_key_from_py(
key: &Bound<'_, PyAny>,
pub fn decoding_key_from_py<'a>(
key: &'a Bound<'_, PyAny>,
algorithms: &[Algorithm],
) -> PyResult<JwtDecodingKey> {
if let Ok(key_ref) = key.extract::<PyRef<'_, DecodingKey>>() {
return key_ref.material.decoding_key(algorithms);
) -> PyResult<BorrowedDecodingKey<'a>> {
if let Ok(bound) = key.cast::<DecodingKey>() {
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<JwtDecodingKey> {
) -> PyResult<BorrowedDecodingKey<'a>> {
let family = ensure_single_family(algorithms)?;
if family != KeyFamily::Hmac {
return Err(errors::invalid_key(
Expand All @@ -296,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.
Expand Down
Loading