Skip to main content

jwtoken/
algorithm.rs

1//! Algorithm implementations for JWT signing and verification.
2//!
3//! This module provides traits and implementations for various JWT signing algorithms.
4//! Supports HMAC-SHA256 (HS256).
5
6use crate::JwtError;
7
8/// Trait for JWT signing algorithms.
9pub trait Signer {
10    /// Returns the algorithm name.
11    fn name(&self) -> &str;
12
13    /// Signs a message and returns the signature.
14    fn sign(&self, message: &[u8]) -> Result<Vec<u8>, JwtError>;
15}
16
17/// Trait for JWT verification algorithms.
18pub trait Verifier {
19    /// Returns the algorithm name.
20    fn name(&self) -> &str;
21
22    /// Verifies a message against a signature.
23    fn verify(&self, message: &[u8], signature: &[u8]) -> Result<bool, JwtError>;
24}
25
26// HMAC-SHA256 implementation (HS256)
27#[cfg(feature = "hs256")]
28mod hs256 {
29    use super::*;
30    use hmac::{Hmac, Mac};
31    use sha2::Sha256;
32
33    /// HMAC-SHA256 (HS256) algorithm implementation.
34    #[derive(Debug, Clone)]
35    pub struct HS256 {
36        secret: Vec<u8>,
37    }
38
39    impl HS256 {
40        /// Creates a new `HS256` instance with the given secret.
41        pub fn new(secret: &[u8]) -> Self {
42            Self {
43                secret: secret.to_vec(),
44            }
45        }
46    }
47
48    impl Signer for HS256 {
49        fn name(&self) -> &str {
50            "HS256"
51        }
52
53        fn sign(&self, message: &[u8]) -> Result<Vec<u8>, JwtError> {
54            let mut mac = Hmac::<Sha256>::new_from_slice(&self.secret)
55                .map_err(|_| JwtError::InvalidSignature)?;
56            mac.update(message);
57
58            Ok(mac.finalize().into_bytes().to_vec())
59        }
60    }
61
62    impl Verifier for HS256 {
63        fn name(&self) -> &str {
64            "HS256"
65        }
66
67        fn verify(&self, message: &[u8], signature: &[u8]) -> Result<bool, JwtError> {
68            let expected = self.sign(message)?;
69            Ok(expected == signature)
70        }
71    }
72}
73
74#[cfg(feature = "rs256")]
75mod rs256 {
76    use super::*;
77    use rsa::pkcs1::{DecodeRsaPrivateKey, DecodeRsaPublicKey};
78    use rsa::pkcs1v15::{SigningKey, VerifyingKey};
79    use rsa::pkcs8::{DecodePrivateKey, DecodePublicKey};
80    use rsa::signature::{RandomizedSigner, SignatureEncoding, Verifier as RsaVerifier};
81    use rsa::{RsaPrivateKey, RsaPublicKey};
82    use sha2::Sha256;
83
84    /// RSA-SHA256 (RS256) signer implementation with private key.
85    #[derive(Debug, Clone)]
86    pub struct RS256Signer {
87        private_key: RsaPrivateKey,
88        public_key: RsaPublicKey, // Cache the public key for verification
89    }
90
91    /// RSA-SHA256 (RS256) verifier implementation with public key.
92    #[derive(Debug, Clone)]
93    pub struct RS256Verifier {
94        public_key: RsaPublicKey,
95    }
96
97    impl RS256Signer {
98        /// Creates a new `RS256Signer` instance with the given private key.
99        pub fn new(private_key: RsaPrivateKey) -> Self {
100            let public_key = private_key.to_public_key();
101            Self {
102                private_key,
103                public_key,
104            }
105        }
106
107        /// Creates a new `RS256Signer` from PEM-encoded private key bytes.
108        pub fn from_pem(pem_bytes: &[u8]) -> Result<Self, JwtError> {
109            let pem_str = std::str::from_utf8(pem_bytes).map_err(|_| JwtError::InvalidKey)?;
110
111            // Try PKCS#8 format first
112            let private_key = RsaPrivateKey::from_pkcs8_pem(pem_str)
113                .or_else(|_| RsaPrivateKey::from_pkcs1_pem(pem_str))
114                .map_err(|_| JwtError::InvalidKey)?;
115
116            Ok(Self::new(private_key))
117        }
118
119        /// Creates a new `RS256Signer` from DER-encoded private key bytes.
120        pub fn from_der(der_bytes: &[u8]) -> Result<Self, JwtError> {
121            // Try PKCS#8 format first
122            let private_key = RsaPrivateKey::from_pkcs8_der(der_bytes)
123                .or_else(|_| RsaPrivateKey::from_pkcs1_der(der_bytes))
124                .map_err(|_| JwtError::InvalidKey)?;
125
126            Ok(Self::new(private_key))
127        }
128
129        /// Get the public key for this signer
130        pub fn public_key(&self) -> &RsaPublicKey {
131            &self.public_key
132        }
133    }
134
135    impl RS256Verifier {
136        /// Creates a new `RS256Verifier` instance with the given public key.
137        pub fn new(public_key: RsaPublicKey) -> Self {
138            Self { public_key }
139        }
140
141        /// Creates a new `RS256Verifier` from PEM-encoded public key bytes.
142        pub fn from_pem(pem_bytes: &[u8]) -> Result<Self, JwtError> {
143            let pem_str = std::str::from_utf8(pem_bytes).map_err(|_| JwtError::InvalidKey)?;
144
145            let public_key = RsaPublicKey::from_public_key_pem(pem_str)
146                .or_else(|_| RsaPublicKey::from_pkcs1_pem(pem_str))
147                .map_err(|_| JwtError::InvalidKey)?;
148
149            Ok(Self::new(public_key))
150        }
151
152        /// Creates a new `RS256Verifier` from DER-encoded public key bytes.
153        pub fn from_der(der_bytes: &[u8]) -> Result<Self, JwtError> {
154            let public_key = RsaPublicKey::from_public_key_der(der_bytes)
155                .or_else(|_| RsaPublicKey::from_pkcs1_der(der_bytes))
156                .map_err(|_| JwtError::InvalidKey)?;
157
158            Ok(Self::new(public_key))
159        }
160    }
161
162    impl Signer for RS256Signer {
163        fn name(&self) -> &str {
164            "RS256"
165        }
166
167        fn sign(&self, message: &[u8]) -> Result<Vec<u8>, JwtError> {
168            let signing_key = SigningKey::<Sha256>::new(self.private_key.clone());
169            let mut rng = rsa::rand_core::OsRng;
170
171            let signature = signing_key.sign_with_rng(&mut rng, message);
172            Ok(signature.to_vec())
173        }
174    }
175
176    impl Verifier for RS256Verifier {
177        fn name(&self) -> &str {
178            "RS256"
179        }
180
181        fn verify(&self, message: &[u8], signature: &[u8]) -> Result<bool, JwtError> {
182            let verifying_key = VerifyingKey::<Sha256>::new(self.public_key.clone());
183            let signature = rsa::pkcs1v15::Signature::try_from(signature)
184                .map_err(|_| JwtError::InvalidSignature)?;
185
186            Ok(verifying_key.verify(message, &signature).is_ok())
187        }
188    }
189
190    impl Verifier for RS256Signer {
191        fn name(&self) -> &str {
192            "RS256"
193        }
194
195        fn verify(&self, message: &[u8], signature: &[u8]) -> Result<bool, JwtError> {
196            let verifying_key = VerifyingKey::<Sha256>::new(self.public_key.clone());
197            let signature = rsa::pkcs1v15::Signature::try_from(signature)
198                .map_err(|_| JwtError::InvalidSignature)?;
199
200            Ok(verifying_key.verify(message, &signature).is_ok())
201        }
202    }
203}
204
205#[cfg(feature = "hs256")]
206pub use hs256::*;
207
208#[cfg(feature = "rs256")]
209pub use rs256::*;