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)]
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 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
98pub fn compute_key_id(spki_bytes: &[u8]) -> String {
102 let hash = Sha256::digest(spki_bytes);
103 format!("sha256:{}", hex::encode(hash))
104}
105
106pub 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
112fn 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
121fn build_pae(payload_type: &str, payload: &[u8]) -> Vec<u8> {
126 assay_common::dsse::build_pae(payload_type, payload)
127}
128
129fn 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
138fn compute_payload_digest(canonical: &[u8]) -> String {
140 let hash = Sha256::digest(canonical);
141 format!("sha256:{}", hex::encode(hash))
142}
143
144pub fn sign_tool(tool: &Value, signing_key: &SigningKey, embed_pubkey: bool) -> Result<Value> {
156 let tool_without_sig = strip_signature(tool)?;
158
159 let canonical = jcs::to_vec(&tool_without_sig)?;
161
162 let pae = build_pae(PAYLOAD_TYPE_TOOL_V1, &canonical);
164
165 let signature: Signature = signing_key.sign(&pae);
167
168 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 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 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
201pub fn verify_tool(tool: &Value, trusted_key: &VerifyingKey) -> Result<VerifyResult, VerifyError> {
212 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 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 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 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 let computed_digest = compute_payload_digest(&canonical);
251 if sig.payload_digest != computed_digest {
252 return Err(VerifyError::DigestMismatch);
253 }
254
255 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 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
293pub 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
299pub 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 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 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); }
397
398 #[test]
399 fn test_pae_format() {
400 let pae = build_pae("application/json", b"test");
401
402 let expected = b"DSSEv1 16 application/json 4 test";
404 assert_eq!(pae, expected);
405 }
406
407 #[test]
413 fn test_payload_type_length_normative() {
414 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 assert!(payload_type.is_ascii());
423
424 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]
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 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 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 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}