Skip to main content

mcp_airlock/
crypto.rs

1//! # DPoP Cryptographic Primitives
2//!
3//! This module provides the implementation of **Demonstrating Proof-of-Possession** (DPoP)
4//! as per [RFC 9449](https://datatracker.ietf.org/doc/html/rfc9449).
5//!
6//! It handles:
7//! - P-256 key pair generation.
8//! - DPoP-signed JWT generation.
9//! - SHA-256 hashing of access tokens for the `ath` claim.
10
11use crate::Result;
12use anyhow::Context;
13use base64::{engine::general_purpose::URL_SAFE_NO_PAD, Engine as _};
14use p256::ecdsa::{signature::Signer, SigningKey, VerifyingKey};
15use p256::SecretKey;
16use rand_core::OsRng;
17use serde::{Deserialize, Serialize};
18use serde_json::{json, Value};
19use sha2::{Digest, Sha256};
20use std::time::{SystemTime, UNIX_EPOCH};
21use uuid::Uuid;
22
23/// An ephemeral P-256 key pair used for signing DPoP proofs.
24pub struct DpopKey {
25    /// The ECDSA signing key.
26    signing_key: SigningKey,
27}
28
29#[derive(Debug, Serialize, Deserialize)]
30struct DpopClaims {
31    jti: String,
32    htm: String,
33    htu: String,
34    iat: u64,
35    #[serde(skip_serializing_if = "Option::is_none")]
36    ath: Option<String>,
37    /// Server-provided nonce (RFC 9449 §8, §9).
38    #[serde(skip_serializing_if = "Option::is_none")]
39    nonce: Option<String>,
40}
41
42impl DpopKey {
43    /// Generates a new ephemeral DPoP keypair.
44    pub fn generate() -> Self {
45        let signing_key = SigningKey::random(&mut OsRng);
46        Self { signing_key }
47    }
48
49    /// Restores a DPoP keypair from raw bytes.
50    pub fn from_bytes(bytes: &[u8]) -> Result<Self> {
51        let secret_key = SecretKey::from_slice(bytes).context("Invalid DPoP key bytes")?;
52        let signing_key = SigningKey::from(secret_key);
53        Ok(Self { signing_key })
54    }
55
56    /// Exports the private key as bytes for secure storage.
57    pub fn to_bytes(&self) -> Vec<u8> {
58        self.signing_key.to_bytes().to_vec()
59    }
60
61    /// Constructs the public JWK representation.
62    pub fn public_jwk(&self) -> Result<Value> {
63        let verifying_key = VerifyingKey::from(&self.signing_key);
64        let encoded_point = verifying_key.to_encoded_point(false);
65
66        Ok(json!({
67            "kty": "EC",
68            "crv": "P-256",
69            "x": URL_SAFE_NO_PAD.encode(encoded_point.x().context("P-256 must have x")?),
70            "y": URL_SAFE_NO_PAD.encode(encoded_point.y().context("P-256 must have y")?),
71        }))
72    }
73
74    /// The RFC 7638 JWK thumbprint of the public key (the DPoP `jkt`).
75    pub fn jkt(&self) -> Result<String> {
76        let jwk = self.public_jwk()?;
77        let coord = |k: &str| {
78            jwk[k]
79                .as_str()
80                .map(str::to_string)
81                .context("missing JWK coordinate")
82        };
83        Ok(ec_thumbprint(&coord("x")?, &coord("y")?))
84    }
85
86    /// Generates a DPoP Proof JWT for a given HTTP method and URL.
87    /// Optional access_token can be provided to include 'ath' claim.
88    pub fn generate_proof(&self, htm: &str, htu: &str) -> Result<String> {
89        self.generate_proof_with_ath(htm, htu, None, None)
90    }
91
92    /// Generates a DPoP Proof JWT, optionally with an access token hash (`ath`)
93    /// and a server-provided `nonce`.
94    pub fn generate_proof_with_ath(
95        &self,
96        htm: &str,
97        htu: &str,
98        access_token: Option<&str>,
99        nonce: Option<&str>,
100    ) -> Result<String> {
101        let jwk = self.public_jwk()?;
102
103        let header = json!({
104            "typ": "dpop+jwt",
105            "alg": "ES256",
106            "jwk": jwk
107        });
108
109        let now = SystemTime::now().duration_since(UNIX_EPOCH)?.as_secs();
110
111        let ath = access_token.map(|at| {
112            let mut hasher = Sha256::new();
113            hasher.update(at.as_bytes());
114            URL_SAFE_NO_PAD.encode(hasher.finalize())
115        });
116
117        let claims = DpopClaims {
118            jti: Uuid::new_v4().to_string(),
119            htm: htm.to_string(),
120            htu: normalize_htu(htu),
121            iat: now,
122            ath,
123            nonce: nonce.map(str::to_string),
124        };
125
126        let header_str = serde_json::to_string(&header)?;
127        let claims_str = serde_json::to_string(&claims)?;
128
129        // Pre-allocate buffer for the JWT message (header + '.' + payload)
130        // Base64 encoding size is roughly 4/3 of the input size
131        let header_len = (header_str.len() * 4).div_ceil(3);
132        let claims_len = (claims_str.len() * 4).div_ceil(3);
133        let mut message = String::with_capacity(header_len + 1 + claims_len);
134
135        URL_SAFE_NO_PAD.encode_string(header_str.as_bytes(), &mut message);
136        message.push('.');
137        URL_SAFE_NO_PAD.encode_string(claims_str.as_bytes(), &mut message);
138
139        let signature: p256::ecdsa::Signature = self.signing_key.sign(message.as_bytes());
140
141        // Pre-allocate buffer for the final JWT (message + '.' + signature)
142        let sig_len = 86; // approximate length of base64url encoded P-256 signature
143        let mut final_jwt = String::with_capacity(message.len() + 1 + sig_len);
144        final_jwt.push_str(&message);
145        final_jwt.push('.');
146        URL_SAFE_NO_PAD.encode_string(signature.to_bytes(), &mut final_jwt);
147
148        Ok(final_jwt)
149    }
150}
151
152/// RFC 7638 thumbprint of a P-256 key: SHA-256 over the required members in
153/// lexicographic order, without whitespace.
154fn ec_thumbprint(x: &str, y: &str) -> String {
155    // 40 is the length of `{"crv":"P-256","kty":"EC","x":"","y":""}`
156    let mut canonical = String::with_capacity(40 + x.len() + y.len());
157    canonical.push_str(r#"{"crv":"P-256","kty":"EC","x":""#);
158    canonical.push_str(x);
159    canonical.push_str(r#"","y":""#);
160    canonical.push_str(y);
161    canonical.push_str(r#""}"#);
162    URL_SAFE_NO_PAD.encode(Sha256::digest(canonical.as_bytes()))
163}
164
165/// The `DPoP-Nonce` response header, if present (RFC 9449 §8.1).
166pub fn dpop_nonce(headers: &reqwest::header::HeaderMap) -> Option<String> {
167    headers
168        .get("DPoP-Nonce")
169        .and_then(|v| v.to_str().ok())
170        .filter(|v| !v.is_empty())
171        .map(str::to_string)
172}
173
174/// The `htu` claim is the target URI without query and fragment (RFC 9449 §4.2).
175fn normalize_htu(htu: &str) -> String {
176    match url::Url::parse(htu) {
177        Ok(mut u) if u.query().is_some() || u.fragment().is_some() => {
178            u.set_query(None);
179            u.set_fragment(None);
180            u.into()
181        }
182        _ => htu.to_string(),
183    }
184}
185
186#[cfg(test)]
187mod tests {
188    use super::*;
189
190    #[test]
191    fn test_ec_thumbprint_rfc9449_example() {
192        // RFC 9449 §6.1 / §10: the example key and its jkt.
193        assert_eq!(
194            ec_thumbprint(
195                "l8tFrhx-34tV3hRICRDY9zCkDlpBhF42UQUfWVAWBFs",
196                "9VE4jf_Ok_o64zbTTlcuNJajHmt6v9TDVrU0CdvGRDA"
197            ),
198            "0ZcOCORZNYy-DWpqq30jZyJGHTN0d2HglBV3uiguA4I"
199        );
200    }
201
202    #[test]
203    fn test_jkt_matches_proof_jwk() -> Result<()> {
204        let key = DpopKey::generate();
205        let proof = key.generate_proof("POST", "https://as/token")?;
206        let header: Value =
207            serde_json::from_slice(&URL_SAFE_NO_PAD.decode(proof.split('.').next().unwrap())?)?;
208        let jwk = &header["jwk"];
209        assert_eq!(
210            key.jkt()?,
211            ec_thumbprint(jwk["x"].as_str().unwrap(), jwk["y"].as_str().unwrap())
212        );
213        assert_eq!(key.jkt()?.len(), 43);
214        Ok(())
215    }
216
217    #[test]
218    fn test_normalize_htu() {
219        assert_eq!(
220            normalize_htu("https://api.example.com/rpc?session=1#frag"),
221            "https://api.example.com/rpc"
222        );
223        assert_eq!(
224            normalize_htu("https://api.example.com/rpc"),
225            "https://api.example.com/rpc"
226        );
227        assert_eq!(normalize_htu("not a url"), "not a url");
228    }
229
230    #[test]
231    fn test_dpop_key_generate_and_bytes() -> Result<()> {
232        let key = DpopKey::generate();
233        let bytes = key.to_bytes();
234        assert_eq!(bytes.len(), 32);
235
236        let key2 = DpopKey::from_bytes(&bytes)?;
237        assert_eq!(key.to_bytes(), key2.to_bytes());
238        Ok(())
239    }
240
241    #[test]
242    fn test_dpop_key_invalid_bytes() {
243        let res = DpopKey::from_bytes(&[1, 2, 3]);
244        assert!(res.is_err());
245        assert!(format!("{:?}", res.err().unwrap()).contains("Invalid DPoP key bytes"));
246    }
247
248    #[test]
249    fn test_public_jwk() -> Result<()> {
250        let key = DpopKey::generate();
251        let jwk = key.public_jwk()?;
252        assert_eq!(jwk["kty"], "EC");
253        assert_eq!(jwk["crv"], "P-256");
254        assert!(jwk.get("x").is_some());
255        assert!(jwk.get("y").is_some());
256        Ok(())
257    }
258
259    #[test]
260    fn test_generate_proof() -> Result<()> {
261        let key = DpopKey::generate();
262        let proof = key.generate_proof("POST", "https://api.example.com/rpc")?;
263        let parts: Vec<&str> = proof.split('.').collect();
264        assert_eq!(parts.len(), 3);
265
266        let header_json: Value = serde_json::from_slice(&URL_SAFE_NO_PAD.decode(parts[0])?)?;
267        assert_eq!(header_json["typ"], "dpop+jwt");
268        assert_eq!(header_json["alg"], "ES256");
269        assert!(header_json.get("jwk").is_some());
270
271        let claims_json: Value = serde_json::from_slice(&URL_SAFE_NO_PAD.decode(parts[1])?)?;
272        assert_eq!(claims_json["htm"], "POST");
273        assert_eq!(claims_json["htu"], "https://api.example.com/rpc");
274        assert!(claims_json.get("jti").is_some());
275        assert!(claims_json.get("iat").is_some());
276        Ok(())
277    }
278
279    #[test]
280    fn test_generate_proof_success() -> Result<()> {
281        let key = DpopKey::generate();
282        let proof = key.generate_proof("GET", "https://api.example.com/resource")?;
283
284        let parts: Vec<&str> = proof.split('.').collect();
285        assert_eq!(parts.len(), 3, "Proof must have 3 parts");
286
287        let claims_json: Value = serde_json::from_slice(&URL_SAFE_NO_PAD.decode(parts[1])?)?;
288        assert_eq!(claims_json["htm"], "GET");
289        assert_eq!(claims_json["htu"], "https://api.example.com/resource");
290
291        // Assert that 'ath' claim is not present when calling generate_proof (wrapper without ath)
292        assert!(
293            claims_json.get("ath").is_none(),
294            "ath claim should not be present in wrapper generate_proof"
295        );
296
297        Ok(())
298    }
299
300    #[test]
301    fn test_generate_proof_signature() -> Result<()> {
302        use p256::ecdsa::signature::Verifier;
303
304        let key = DpopKey::generate();
305        let htm = "POST";
306        let htu = "https://api.example.com/rpc";
307        let proof = key.generate_proof(htm, htu)?;
308        let parts: Vec<&str> = proof.split('.').collect();
309        assert_eq!(parts.len(), 3);
310
311        let message = format!("{}.{}", parts[0], parts[1]);
312        let sig_bytes = URL_SAFE_NO_PAD.decode(parts[2])?;
313        let signature = p256::ecdsa::Signature::from_slice(&sig_bytes)?;
314
315        let verifying_key = VerifyingKey::from(&key.signing_key);
316        verifying_key
317            .verify(message.as_bytes(), &signature)
318            .expect("Signature verification failed");
319
320        let claims_json: Value = serde_json::from_slice(&URL_SAFE_NO_PAD.decode(parts[1])?)?;
321        assert_eq!(claims_json["htm"], htm);
322        assert_eq!(claims_json["htu"], htu);
323
324        Ok(())
325    }
326
327    #[test]
328    fn test_generate_proof_with_nonce() -> Result<()> {
329        let key = DpopKey::generate();
330        let decode = |proof: String| -> Result<Value> {
331            let claims = proof.split('.').nth(1).unwrap().to_string();
332            Ok(serde_json::from_slice(&URL_SAFE_NO_PAD.decode(claims)?)?)
333        };
334
335        let with =
336            decode(key.generate_proof_with_ath("POST", "https://as/token", None, Some("n-1"))?)?;
337        assert_eq!(with["nonce"], "n-1");
338        let without = decode(key.generate_proof("POST", "https://as/token")?)?;
339        assert!(without.get("nonce").is_none());
340        Ok(())
341    }
342
343    #[test]
344    fn test_generate_proof_with_ath() -> Result<()> {
345        let key = DpopKey::generate();
346        let access_token = "test_token";
347        let proof = key.generate_proof_with_ath(
348            "GET",
349            "https://api.example.com/sse",
350            Some(access_token),
351            None,
352        )?;
353        let parts: Vec<&str> = proof.split('.').collect();
354        assert_eq!(parts.len(), 3);
355
356        let header_json: Value = serde_json::from_slice(&URL_SAFE_NO_PAD.decode(parts[0])?)?;
357        assert_eq!(header_json["typ"], "dpop+jwt");
358        assert_eq!(header_json["alg"], "ES256");
359        assert!(header_json.get("jwk").is_some());
360
361        let claims_json: Value = serde_json::from_slice(&URL_SAFE_NO_PAD.decode(parts[1])?)?;
362        assert_eq!(claims_json["htm"], "GET");
363        assert_eq!(claims_json["htu"], "https://api.example.com/sse");
364        assert!(claims_json.get("jti").is_some());
365        assert!(claims_json.get("iat").is_some());
366        assert!(claims_json.get("ath").is_some());
367
368        let mut hasher = Sha256::new();
369        hasher.update(access_token.as_bytes());
370        let expected_ath = URL_SAFE_NO_PAD.encode(hasher.finalize());
371        assert_eq!(claims_json["ath"], expected_ath);
372        Ok(())
373    }
374}