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