Skip to main content

assay_core/mcp/
signing.rs

1//! Tool signing and verification per SPEC-Tool-Signing-v1.
2//!
3//! Provides ed25519 signing/verification with DSSE-compatible PAE encoding.
4
5use anyhow::{bail, Context, Result};
6use base64::{engine::general_purpose::STANDARD as BASE64, Engine};
7use chrono::{DateTime, Utc};
8use ed25519_dalek::{Signature, Signer, SigningKey, Verifier, VerifyingKey};
9use serde::{Deserialize, Serialize};
10use serde_json::Value;
11use sha2::{Digest, Sha256};
12
13use super::jcs;
14
15/// Payload type for tool definitions (DSSE-style binding).
16pub const PAYLOAD_TYPE_TOOL_V1: &str = "application/vnd.assay.tool+json;v=1";
17
18/// The x-assay-sig field name.
19pub const SIG_FIELD: &str = "x-assay-sig";
20
21/// Signature algorithm.
22#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
23#[serde(rename_all = "lowercase")]
24pub enum SignatureAlgorithm {
25    Ed25519,
26}
27
28/// The x-assay-sig structure.
29///
30/// # Field Serialization
31///
32/// Producers SHOULD omit `public_key` when not embedding the key.
33/// This is enforced via `skip_serializing_if = "Option::is_none"`.
34/// Verifiers MUST treat `null` as equivalent to absent.
35#[derive(Debug, Clone, Serialize, Deserialize)]
36pub struct ToolSignature {
37    pub version: u8,
38    pub algorithm: SignatureAlgorithm,
39    pub payload_type: String,
40    pub payload_digest: String,
41    pub key_id: String,
42    pub signature: String,
43    pub signed_at: DateTime<Utc>,
44    /// Embedded public key (SPKI DER, base64).
45    /// Producers SHOULD omit (not set to null) when not embedding.
46    #[serde(default, skip_serializing_if = "Option::is_none")]
47    pub public_key: Option<String>,
48}
49
50/// Result of successful verification.
51#[derive(Debug, Clone)]
52pub struct VerifyResult {
53    pub key_id: String,
54    pub signed_at: DateTime<Utc>,
55}
56
57/// Verification errors with exit codes.
58#[derive(Debug, Clone, thiserror::Error)]
59#[non_exhaustive]
60pub enum VerifyError {
61    #[error("tool is not signed")]
62    NoSignature,
63
64    #[error("payload type mismatch: expected {expected}, got {got}")]
65    PayloadTypeMismatch { expected: String, got: String },
66
67    #[error("signature invalid: {reason}")]
68    SignatureInvalid { reason: String },
69
70    #[error("key not trusted: {key_id}")]
71    KeyNotTrusted { key_id: String },
72
73    #[error("malformed signature: {reason}")]
74    MalformedSignature { reason: String },
75
76    #[error("payload digest mismatch")]
77    DigestMismatch,
78
79    #[error("key_id mismatch: signature claims {claimed}, actual {actual}")]
80    KeyIdMismatch { claimed: String, actual: String },
81}
82
83impl VerifyError {
84    /// Exit code for CLI.
85    pub fn exit_code(&self) -> i32 {
86        match self {
87            Self::NoSignature => 2,
88            Self::KeyNotTrusted { .. } => 3,
89            Self::SignatureInvalid { .. }
90            | Self::PayloadTypeMismatch { .. }
91            | Self::DigestMismatch
92            | Self::KeyIdMismatch { .. } => 4,
93            Self::MalformedSignature { .. } => 1,
94        }
95    }
96}
97
98/// Compute key_id from SPKI-encoded public key bytes.
99///
100/// Returns `sha256:<lowercase-hex>`.
101pub fn compute_key_id(spki_bytes: &[u8]) -> String {
102    let hash = Sha256::digest(spki_bytes);
103    format!("sha256:{}", hex::encode(hash))
104}
105
106/// Compute key_id from a VerifyingKey.
107pub fn compute_key_id_from_verifying_key(key: &VerifyingKey) -> Result<String> {
108    let spki_bytes = key_to_spki_der(key)?;
109    Ok(compute_key_id(&spki_bytes))
110}
111
112/// Convert VerifyingKey to SPKI DER bytes.
113fn key_to_spki_der(key: &VerifyingKey) -> Result<Vec<u8>> {
114    use ed25519_dalek::pkcs8::EncodePublicKey;
115    let doc = key
116        .to_public_key_der()
117        .context("failed to encode public key as SPKI DER")?;
118    Ok(doc.as_bytes().to_vec())
119}
120
121/// Build DSSE Pre-Authentication Encoding (PAE).
122///
123/// Delegates to [`assay_common::dsse::build_pae`], which `assay-evidence` also
124/// calls. PAE defines what a signature covers, so one construction serves both.
125fn build_pae(payload_type: &str, payload: &[u8]) -> Vec<u8> {
126    assay_common::dsse::build_pae(payload_type, payload)
127}
128
129/// Remove x-assay-sig field from tool JSON.
130fn strip_signature(tool: &Value) -> Result<Value> {
131    let mut tool = tool.clone();
132    if let Some(obj) = tool.as_object_mut() {
133        obj.remove(SIG_FIELD);
134    }
135    Ok(tool)
136}
137
138/// Compute payload digest.
139fn compute_payload_digest(canonical: &[u8]) -> String {
140    let hash = Sha256::digest(canonical);
141    format!("sha256:{}", hex::encode(hash))
142}
143
144/// Sign a tool definition.
145///
146/// # Arguments
147///
148/// * `tool` - Tool definition JSON (may or may not have existing signature)
149/// * `signing_key` - Ed25519 private key
150/// * `embed_pubkey` - If true, include public_key in signature
151///
152/// # Returns
153///
154/// Tool definition with x-assay-sig field added.
155pub fn sign_tool(tool: &Value, signing_key: &SigningKey, embed_pubkey: bool) -> Result<Value> {
156    // 1. Remove existing signature
157    let tool_without_sig = strip_signature(tool)?;
158
159    // 2. Canonicalize
160    let canonical = jcs::to_vec(&tool_without_sig)?;
161
162    // 3. Build PAE
163    let pae = build_pae(PAYLOAD_TYPE_TOOL_V1, &canonical);
164
165    // 4. Sign
166    let signature: Signature = signing_key.sign(&pae);
167
168    // 5. Compute digests
169    let payload_digest = compute_payload_digest(&canonical);
170    let verifying_key = signing_key.verifying_key();
171    let key_id = compute_key_id_from_verifying_key(&verifying_key)?;
172
173    // 6. Build x-assay-sig
174    let sig = ToolSignature {
175        version: 1,
176        algorithm: SignatureAlgorithm::Ed25519,
177        payload_type: PAYLOAD_TYPE_TOOL_V1.to_string(),
178        payload_digest,
179        key_id,
180        signature: BASE64.encode(signature.to_bytes()),
181        signed_at: Utc::now(),
182        public_key: if embed_pubkey {
183            let spki = key_to_spki_der(&verifying_key)?;
184            Some(BASE64.encode(&spki))
185        } else {
186            None
187        },
188    };
189
190    // 7. Add to tool
191    let mut result = tool_without_sig;
192    if let Some(obj) = result.as_object_mut() {
193        obj.insert(SIG_FIELD.to_string(), serde_json::to_value(&sig)?);
194    } else {
195        bail!("tool must be a JSON object");
196    }
197
198    Ok(result)
199}
200
201/// Verify a signed tool definition.
202///
203/// # Arguments
204///
205/// * `tool` - Signed tool definition JSON
206/// * `trusted_key` - Public key to verify against
207///
208/// # Returns
209///
210/// `VerifyResult` on success, `VerifyError` on failure.
211pub fn verify_tool(tool: &Value, trusted_key: &VerifyingKey) -> Result<VerifyResult, VerifyError> {
212    // 1. Extract signature
213    let sig_value = tool.get(SIG_FIELD).ok_or(VerifyError::NoSignature)?;
214
215    let sig: ToolSignature =
216        serde_json::from_value(sig_value.clone()).map_err(|e| VerifyError::MalformedSignature {
217            reason: e.to_string(),
218        })?;
219
220    // 2. Validate version and algorithm
221    if sig.version != 1 {
222        return Err(VerifyError::MalformedSignature {
223            reason: format!("unsupported version: {}", sig.version),
224        });
225    }
226    if sig.algorithm != SignatureAlgorithm::Ed25519 {
227        return Err(VerifyError::MalformedSignature {
228            reason: format!("unsupported algorithm: {:?}", sig.algorithm),
229        });
230    }
231
232    // 3. Validate payload_type
233    if sig.payload_type != PAYLOAD_TYPE_TOOL_V1 {
234        return Err(VerifyError::PayloadTypeMismatch {
235            expected: PAYLOAD_TYPE_TOOL_V1.to_string(),
236            got: sig.payload_type,
237        });
238    }
239
240    // 4. Strip signature and canonicalize
241    let tool_without_sig = strip_signature(tool).map_err(|e| VerifyError::MalformedSignature {
242        reason: e.to_string(),
243    })?;
244    let canonical =
245        jcs::to_vec(&tool_without_sig).map_err(|e| VerifyError::MalformedSignature {
246            reason: e.to_string(),
247        })?;
248
249    // 5. Verify payload digest
250    let computed_digest = compute_payload_digest(&canonical);
251    if sig.payload_digest != computed_digest {
252        return Err(VerifyError::DigestMismatch);
253    }
254
255    // 6. Build PAE and verify signature
256    let pae = build_pae(&sig.payload_type, &canonical);
257    let signature_bytes =
258        BASE64
259            .decode(&sig.signature)
260            .map_err(|e| VerifyError::MalformedSignature {
261                reason: format!("invalid base64 signature: {}", e),
262            })?;
263    let signature =
264        Signature::from_slice(&signature_bytes).map_err(|e| VerifyError::MalformedSignature {
265            reason: format!("invalid signature bytes: {}", e),
266        })?;
267
268    trusted_key
269        .verify(&pae, &signature)
270        .map_err(|_| VerifyError::SignatureInvalid {
271            reason: "ed25519 verification failed".to_string(),
272        })?;
273
274    // 7. Verify key_id matches
275    let actual_key_id = compute_key_id_from_verifying_key(trusted_key).map_err(|e| {
276        VerifyError::MalformedSignature {
277            reason: e.to_string(),
278        }
279    })?;
280    if sig.key_id != actual_key_id {
281        return Err(VerifyError::KeyIdMismatch {
282            claimed: sig.key_id,
283            actual: actual_key_id,
284        });
285    }
286
287    Ok(VerifyResult {
288        key_id: sig.key_id,
289        signed_at: sig.signed_at,
290    })
291}
292
293/// Extract signature from a tool (if present).
294pub fn extract_signature(tool: &Value) -> Option<ToolSignature> {
295    tool.get(SIG_FIELD)
296        .and_then(|v| serde_json::from_value(v.clone()).ok())
297}
298
299/// Check if a tool is signed.
300pub fn is_signed(tool: &Value) -> bool {
301    tool.get(SIG_FIELD).is_some()
302}
303
304#[cfg(test)]
305mod tests {
306    use super::*;
307    use getrandom::{rand_core::UnwrapErr, SysRng};
308    use serde_json::json;
309
310    fn generate_keypair() -> SigningKey {
311        let mut csprng = UnwrapErr(SysRng);
312        SigningKey::generate(&mut csprng)
313    }
314
315    #[test]
316    fn test_sign_and_verify_roundtrip() {
317        let key = generate_keypair();
318        let tool = json!({
319            "name": "read_file",
320            "description": "Read a file",
321            "inputSchema": {"type": "object"}
322        });
323
324        let signed = sign_tool(&tool, &key, false).unwrap();
325        assert!(is_signed(&signed));
326
327        let result = verify_tool(&signed, &key.verifying_key()).unwrap();
328        assert!(result.key_id.starts_with("sha256:"));
329    }
330
331    #[test]
332    fn test_tamper_detection() {
333        let key = generate_keypair();
334        let tool = json!({
335            "name": "read_file",
336            "description": "Read a file",
337            "inputSchema": {"type": "object"}
338        });
339
340        let mut signed = sign_tool(&tool, &key, false).unwrap();
341
342        // Tamper with the tool
343        signed["description"] = json!("Malicious description");
344
345        let result = verify_tool(&signed, &key.verifying_key());
346        assert!(matches!(result, Err(VerifyError::DigestMismatch)));
347    }
348
349    #[test]
350    fn test_wrong_key_fails() {
351        let key1 = generate_keypair();
352        let key2 = generate_keypair();
353        let tool = json!({
354            "name": "test_tool",
355            "description": "Test",
356            "inputSchema": {}
357        });
358
359        let signed = sign_tool(&tool, &key1, false).unwrap();
360        let result = verify_tool(&signed, &key2.verifying_key());
361
362        // Should fail with either SignatureInvalid or KeyIdMismatch
363        assert!(matches!(
364            result,
365            Err(VerifyError::SignatureInvalid { .. }) | Err(VerifyError::KeyIdMismatch { .. })
366        ));
367    }
368
369    #[test]
370    fn test_unsigned_tool() {
371        let key = generate_keypair();
372        let tool = json!({"name": "unsigned"});
373
374        let result = verify_tool(&tool, &key.verifying_key());
375        assert!(matches!(result, Err(VerifyError::NoSignature)));
376    }
377
378    #[test]
379    fn test_embed_pubkey() {
380        let key = generate_keypair();
381        let tool = json!({"name": "test", "description": "test", "inputSchema": {}});
382
383        let signed = sign_tool(&tool, &key, true).unwrap();
384        let sig = extract_signature(&signed).unwrap();
385
386        assert!(sig.public_key.is_some());
387    }
388
389    #[test]
390    fn test_key_id_computation() {
391        let key = generate_keypair();
392        let key_id = compute_key_id_from_verifying_key(&key.verifying_key()).unwrap();
393
394        assert!(key_id.starts_with("sha256:"));
395        assert_eq!(key_id.len(), 7 + 64); // "sha256:" + 64 hex chars
396    }
397
398    #[test]
399    fn test_pae_format() {
400        let pae = build_pae("application/json", b"test");
401
402        // "DSSEv1 16 application/json 4 test"
403        let expected = b"DSSEv1 16 application/json 4 test";
404        assert_eq!(pae, expected);
405    }
406
407    /// Normative test vector for PAYLOAD_TYPE_TOOL_V1 length.
408    ///
409    /// This test ensures the exact byte length of the payload type is
410    /// consistent across implementations. PAE uses decimal length encoding,
411    /// so any mismatch causes cross-impl verification failures.
412    #[test]
413    fn test_payload_type_length_normative() {
414        // "application/vnd.assay.tool+json;v=1" is exactly 35 bytes UTF-8
415        let payload_type = PAYLOAD_TYPE_TOOL_V1;
416        assert_eq!(
417            payload_type.len(),
418            35,
419            "PAYLOAD_TYPE_TOOL_V1 must be 35 bytes"
420        );
421        // Verify it's pure ASCII (each char = 1 byte)
422        assert!(payload_type.is_ascii());
423
424        // Verify PAE encoding uses correct length
425        let pae = build_pae(payload_type, b"{}");
426        let pae_str = String::from_utf8_lossy(&pae);
427        assert!(
428            pae_str.starts_with("DSSEv1 35 application/vnd.assay.tool+json;v=1 2 {}"),
429            "PAE must start with 'DSSEv1 35 ...' for tool signing"
430        );
431    }
432
433    /// Test that key_id uses lowercase hex (normative).
434    #[test]
435    fn test_key_id_lowercase_hex() {
436        let key = generate_keypair();
437        let key_id = compute_key_id_from_verifying_key(&key.verifying_key()).unwrap();
438
439        // Must be lowercase hex
440        assert!(key_id.starts_with("sha256:"));
441        let hex_part = &key_id[7..];
442        assert!(
443            hex_part
444                .chars()
445                .all(|c| c.is_ascii_hexdigit() && !c.is_ascii_uppercase()),
446            "key_id hex must be lowercase: {}",
447            key_id
448        );
449    }
450
451    #[test]
452    fn test_canonicalization_stability() {
453        let key = generate_keypair();
454
455        // Same tool, different JSON formatting
456        let tool1 =
457            json!({"name": "test", "description": "desc", "inputSchema": {"type": "object"}});
458        let tool2 =
459            json!({"inputSchema": {"type": "object"}, "name": "test", "description": "desc"});
460
461        let signed1 = sign_tool(&tool1, &key, false).unwrap();
462        let signed2 = sign_tool(&tool2, &key, false).unwrap();
463
464        // Both should have the same payload_digest
465        let sig1 = extract_signature(&signed1).unwrap();
466        let sig2 = extract_signature(&signed2).unwrap();
467
468        assert_eq!(sig1.payload_digest, sig2.payload_digest);
469    }
470
471    #[test]
472    fn test_exit_codes() {
473        assert_eq!(VerifyError::NoSignature.exit_code(), 2);
474        assert_eq!(
475            VerifyError::KeyNotTrusted { key_id: "x".into() }.exit_code(),
476            3
477        );
478        assert_eq!(
479            VerifyError::SignatureInvalid { reason: "x".into() }.exit_code(),
480            4
481        );
482        assert_eq!(
483            VerifyError::MalformedSignature { reason: "x".into() }.exit_code(),
484            1
485        );
486    }
487}