1use 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
15pub const PAYLOAD_TYPE_TOOL_V1: &str = "application/vnd.assay.tool+json;v=1";
17
18pub const SIG_FIELD: &str = "x-assay-sig";
20
21#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
23#[serde(rename_all = "lowercase")]
24pub enum SignatureAlgorithm {
25 Ed25519,
26}
27
28#[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 #[serde(default, skip_serializing_if = "Option::is_none")]
47 pub public_key: Option<String>,
48}
49
50#[derive(Debug, Clone)]
52pub struct VerifyResult {
53 pub key_id: String,
54 pub signed_at: DateTime<Utc>,
55}
56
57#[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 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
97pub fn compute_key_id(spki_bytes: &[u8]) -> String {
101 let hash = Sha256::digest(spki_bytes);
102 format!("sha256:{}", hex::encode(hash))
103}
104
105pub 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
111fn 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
120fn build_pae(payload_type: &str, payload: &[u8]) -> Vec<u8> {
126 let type_len = payload_type.len().to_string();
127 let payload_len = payload.len().to_string();
128
129 let mut pae = Vec::new();
130 pae.extend_from_slice(b"DSSEv1 ");
131 pae.extend_from_slice(type_len.as_bytes());
132 pae.push(b' ');
133 pae.extend_from_slice(payload_type.as_bytes());
134 pae.push(b' ');
135 pae.extend_from_slice(payload_len.as_bytes());
136 pae.push(b' ');
137 pae.extend_from_slice(payload);
138 pae
139}
140
141fn strip_signature(tool: &Value) -> Result<Value> {
143 let mut tool = tool.clone();
144 if let Some(obj) = tool.as_object_mut() {
145 obj.remove(SIG_FIELD);
146 }
147 Ok(tool)
148}
149
150fn compute_payload_digest(canonical: &[u8]) -> String {
152 let hash = Sha256::digest(canonical);
153 format!("sha256:{}", hex::encode(hash))
154}
155
156pub fn sign_tool(tool: &Value, signing_key: &SigningKey, embed_pubkey: bool) -> Result<Value> {
168 let tool_without_sig = strip_signature(tool)?;
170
171 let canonical = jcs::to_vec(&tool_without_sig)?;
173
174 let pae = build_pae(PAYLOAD_TYPE_TOOL_V1, &canonical);
176
177 let signature: Signature = signing_key.sign(&pae);
179
180 let payload_digest = compute_payload_digest(&canonical);
182 let verifying_key = signing_key.verifying_key();
183 let key_id = compute_key_id_from_verifying_key(&verifying_key)?;
184
185 let sig = ToolSignature {
187 version: 1,
188 algorithm: SignatureAlgorithm::Ed25519,
189 payload_type: PAYLOAD_TYPE_TOOL_V1.to_string(),
190 payload_digest,
191 key_id,
192 signature: BASE64.encode(signature.to_bytes()),
193 signed_at: Utc::now(),
194 public_key: if embed_pubkey {
195 let spki = key_to_spki_der(&verifying_key)?;
196 Some(BASE64.encode(&spki))
197 } else {
198 None
199 },
200 };
201
202 let mut result = tool_without_sig;
204 if let Some(obj) = result.as_object_mut() {
205 obj.insert(SIG_FIELD.to_string(), serde_json::to_value(&sig)?);
206 } else {
207 bail!("tool must be a JSON object");
208 }
209
210 Ok(result)
211}
212
213pub fn verify_tool(tool: &Value, trusted_key: &VerifyingKey) -> Result<VerifyResult, VerifyError> {
224 let sig_value = tool.get(SIG_FIELD).ok_or(VerifyError::NoSignature)?;
226
227 let sig: ToolSignature =
228 serde_json::from_value(sig_value.clone()).map_err(|e| VerifyError::MalformedSignature {
229 reason: e.to_string(),
230 })?;
231
232 if sig.version != 1 {
234 return Err(VerifyError::MalformedSignature {
235 reason: format!("unsupported version: {}", sig.version),
236 });
237 }
238 if sig.algorithm != SignatureAlgorithm::Ed25519 {
239 return Err(VerifyError::MalformedSignature {
240 reason: format!("unsupported algorithm: {:?}", sig.algorithm),
241 });
242 }
243
244 if sig.payload_type != PAYLOAD_TYPE_TOOL_V1 {
246 return Err(VerifyError::PayloadTypeMismatch {
247 expected: PAYLOAD_TYPE_TOOL_V1.to_string(),
248 got: sig.payload_type,
249 });
250 }
251
252 let tool_without_sig = strip_signature(tool).map_err(|e| VerifyError::MalformedSignature {
254 reason: e.to_string(),
255 })?;
256 let canonical =
257 jcs::to_vec(&tool_without_sig).map_err(|e| VerifyError::MalformedSignature {
258 reason: e.to_string(),
259 })?;
260
261 let computed_digest = compute_payload_digest(&canonical);
263 if sig.payload_digest != computed_digest {
264 return Err(VerifyError::DigestMismatch);
265 }
266
267 let pae = build_pae(&sig.payload_type, &canonical);
269 let signature_bytes =
270 BASE64
271 .decode(&sig.signature)
272 .map_err(|e| VerifyError::MalformedSignature {
273 reason: format!("invalid base64 signature: {}", e),
274 })?;
275 let signature =
276 Signature::from_slice(&signature_bytes).map_err(|e| VerifyError::MalformedSignature {
277 reason: format!("invalid signature bytes: {}", e),
278 })?;
279
280 trusted_key
281 .verify(&pae, &signature)
282 .map_err(|_| VerifyError::SignatureInvalid {
283 reason: "ed25519 verification failed".to_string(),
284 })?;
285
286 let actual_key_id = compute_key_id_from_verifying_key(trusted_key).map_err(|e| {
288 VerifyError::MalformedSignature {
289 reason: e.to_string(),
290 }
291 })?;
292 if sig.key_id != actual_key_id {
293 return Err(VerifyError::KeyIdMismatch {
294 claimed: sig.key_id,
295 actual: actual_key_id,
296 });
297 }
298
299 Ok(VerifyResult {
300 key_id: sig.key_id,
301 signed_at: sig.signed_at,
302 })
303}
304
305pub fn extract_signature(tool: &Value) -> Option<ToolSignature> {
307 tool.get(SIG_FIELD)
308 .and_then(|v| serde_json::from_value(v.clone()).ok())
309}
310
311pub fn is_signed(tool: &Value) -> bool {
313 tool.get(SIG_FIELD).is_some()
314}
315
316#[cfg(test)]
317mod tests {
318 use super::*;
319 use getrandom::{rand_core::UnwrapErr, SysRng};
320 use serde_json::json;
321
322 fn generate_keypair() -> SigningKey {
323 let mut csprng = UnwrapErr(SysRng);
324 SigningKey::generate(&mut csprng)
325 }
326
327 #[test]
328 fn test_sign_and_verify_roundtrip() {
329 let key = generate_keypair();
330 let tool = json!({
331 "name": "read_file",
332 "description": "Read a file",
333 "inputSchema": {"type": "object"}
334 });
335
336 let signed = sign_tool(&tool, &key, false).unwrap();
337 assert!(is_signed(&signed));
338
339 let result = verify_tool(&signed, &key.verifying_key()).unwrap();
340 assert!(result.key_id.starts_with("sha256:"));
341 }
342
343 #[test]
344 fn test_tamper_detection() {
345 let key = generate_keypair();
346 let tool = json!({
347 "name": "read_file",
348 "description": "Read a file",
349 "inputSchema": {"type": "object"}
350 });
351
352 let mut signed = sign_tool(&tool, &key, false).unwrap();
353
354 signed["description"] = json!("Malicious description");
356
357 let result = verify_tool(&signed, &key.verifying_key());
358 assert!(matches!(result, Err(VerifyError::DigestMismatch)));
359 }
360
361 #[test]
362 fn test_wrong_key_fails() {
363 let key1 = generate_keypair();
364 let key2 = generate_keypair();
365 let tool = json!({
366 "name": "test_tool",
367 "description": "Test",
368 "inputSchema": {}
369 });
370
371 let signed = sign_tool(&tool, &key1, false).unwrap();
372 let result = verify_tool(&signed, &key2.verifying_key());
373
374 assert!(matches!(
376 result,
377 Err(VerifyError::SignatureInvalid { .. }) | Err(VerifyError::KeyIdMismatch { .. })
378 ));
379 }
380
381 #[test]
382 fn test_unsigned_tool() {
383 let key = generate_keypair();
384 let tool = json!({"name": "unsigned"});
385
386 let result = verify_tool(&tool, &key.verifying_key());
387 assert!(matches!(result, Err(VerifyError::NoSignature)));
388 }
389
390 #[test]
391 fn test_embed_pubkey() {
392 let key = generate_keypair();
393 let tool = json!({"name": "test", "description": "test", "inputSchema": {}});
394
395 let signed = sign_tool(&tool, &key, true).unwrap();
396 let sig = extract_signature(&signed).unwrap();
397
398 assert!(sig.public_key.is_some());
399 }
400
401 #[test]
402 fn test_key_id_computation() {
403 let key = generate_keypair();
404 let key_id = compute_key_id_from_verifying_key(&key.verifying_key()).unwrap();
405
406 assert!(key_id.starts_with("sha256:"));
407 assert_eq!(key_id.len(), 7 + 64); }
409
410 #[test]
411 fn test_pae_format() {
412 let pae = build_pae("application/json", b"test");
413
414 let expected = b"DSSEv1 16 application/json 4 test";
416 assert_eq!(pae, expected);
417 }
418
419 #[test]
425 fn test_payload_type_length_normative() {
426 let payload_type = PAYLOAD_TYPE_TOOL_V1;
428 assert_eq!(
429 payload_type.len(),
430 35,
431 "PAYLOAD_TYPE_TOOL_V1 must be 35 bytes"
432 );
433 assert!(payload_type.is_ascii());
435
436 let pae = build_pae(payload_type, b"{}");
438 let pae_str = String::from_utf8_lossy(&pae);
439 assert!(
440 pae_str.starts_with("DSSEv1 35 application/vnd.assay.tool+json;v=1 2 {}"),
441 "PAE must start with 'DSSEv1 35 ...' for tool signing"
442 );
443 }
444
445 #[test]
447 fn test_key_id_lowercase_hex() {
448 let key = generate_keypair();
449 let key_id = compute_key_id_from_verifying_key(&key.verifying_key()).unwrap();
450
451 assert!(key_id.starts_with("sha256:"));
453 let hex_part = &key_id[7..];
454 assert!(
455 hex_part
456 .chars()
457 .all(|c| c.is_ascii_hexdigit() && !c.is_ascii_uppercase()),
458 "key_id hex must be lowercase: {}",
459 key_id
460 );
461 }
462
463 #[test]
464 fn test_canonicalization_stability() {
465 let key = generate_keypair();
466
467 let tool1 =
469 json!({"name": "test", "description": "desc", "inputSchema": {"type": "object"}});
470 let tool2 =
471 json!({"inputSchema": {"type": "object"}, "name": "test", "description": "desc"});
472
473 let signed1 = sign_tool(&tool1, &key, false).unwrap();
474 let signed2 = sign_tool(&tool2, &key, false).unwrap();
475
476 let sig1 = extract_signature(&signed1).unwrap();
478 let sig2 = extract_signature(&signed2).unwrap();
479
480 assert_eq!(sig1.payload_digest, sig2.payload_digest);
481 }
482
483 #[test]
484 fn test_exit_codes() {
485 assert_eq!(VerifyError::NoSignature.exit_code(), 2);
486 assert_eq!(
487 VerifyError::KeyNotTrusted { key_id: "x".into() }.exit_code(),
488 3
489 );
490 assert_eq!(
491 VerifyError::SignatureInvalid { reason: "x".into() }.exit_code(),
492 4
493 );
494 assert_eq!(
495 VerifyError::MalformedSignature { reason: "x".into() }.exit_code(),
496 1
497 );
498 }
499}