protocol.rs (22783B)
1 use std::{ 2 collections::{BTreeMap, BTreeSet}, 3 fs, 4 io::{ErrorKind, Write}, 5 path::{Component, Path, PathBuf}, 6 }; 7 8 use quote::ToTokens; 9 use serde::{Deserialize, Serialize}; 10 use sha2::{Digest, Sha256}; 11 use syn::{Item, Type, Visibility}; 12 use tempfile::NamedTempFile; 13 14 const CONFIG_PATH: &str = "contracts/codegen/protocol_v1.toml"; 15 const PACKAGE: &str = "radroots_protocol"; 16 const HASH_ALGORITHM: &str = "sha256_bytes_v1"; 17 const EXPECTED_SOURCES: &[(&str, &str)] = &[ 18 ("capability::v1", "crates/protocol/src/capability/v1.rs"), 19 ("error::v1", "crates/protocol/src/error/v1.rs"), 20 ("event::v1", "crates/protocol/src/event/v1.rs"), 21 ( 22 "radrootsd::transport_publish::v5", 23 "crates/protocol/src/radrootsd/transport_publish/v5.rs", 24 ), 25 ("runtime::v1", "crates/protocol/src/runtime/v1.rs"), 26 ]; 27 const EXPECTED_MACRO_GENERATED_TYPES: &[(&str, &str, &str, &str, &str)] = &[( 28 "runtime::v1", 29 "crates/protocol/src/runtime/v1.rs", 30 "operation_ids", 31 "OperationId", 32 "enum", 33 )]; 34 35 #[derive(Debug, Deserialize)] 36 #[serde(deny_unknown_fields)] 37 struct Config { 38 schema_version: u16, 39 package: String, 40 source_hash_algorithm: String, 41 inventory_path: String, 42 inventory_sha256_path: String, 43 packaged_inventory_path: String, 44 source: Vec<SourceConfig>, 45 #[serde(default)] 46 macro_generated_type: Vec<MacroGeneratedTypeConfig>, 47 } 48 49 #[derive(Clone, Debug, Deserialize, Eq, Ord, PartialEq, PartialOrd)] 50 #[serde(deny_unknown_fields)] 51 struct SourceConfig { 52 module: String, 53 path: String, 54 } 55 56 #[derive(Clone, Debug, Deserialize, Eq, Ord, PartialEq, PartialOrd)] 57 #[serde(deny_unknown_fields)] 58 struct MacroGeneratedTypeConfig { 59 module: String, 60 path: String, 61 macro_name: String, 62 rust_name: String, 63 kind: String, 64 } 65 66 #[derive(Debug, Serialize)] 67 struct Inventory { 68 schema_version: u16, 69 generator: &'static str, 70 package: &'static str, 71 source_hash_algorithm: &'static str, 72 sources: Vec<SourceInventory>, 73 schemas: Vec<SchemaInventory>, 74 } 75 76 #[derive(Debug, Serialize)] 77 struct SourceInventory { 78 module: String, 79 path: String, 80 sha256: String, 81 types: Vec<TypeInventory>, 82 } 83 84 #[derive(Debug, Serialize)] 85 struct TypeInventory { 86 rust_path: String, 87 kind: &'static str, 88 } 89 90 #[derive(Clone, Debug, Serialize)] 91 struct SchemaInventory { 92 schema_id: String, 93 module: String, 94 generation: u16, 95 } 96 97 struct GeneratedFile { 98 path: PathBuf, 99 display_path: String, 100 bytes: Vec<u8>, 101 } 102 103 pub(crate) fn run(mode: &str, workspace_root: &Path) -> Result<(), String> { 104 let config = load_config(workspace_root)?; 105 validate_config(&config)?; 106 let schemas = protocol_schemas()?; 107 let generated = render_outputs(workspace_root, &config, schemas)?; 108 match mode { 109 "--check" => check_outputs(&generated), 110 "--write" => write_outputs(&generated), 111 _ => Err("usage: cargo xtask generate protocol --check|--write".to_owned()), 112 } 113 } 114 115 pub(crate) fn check(workspace_root: &Path) -> Result<(), String> { 116 run("--check", workspace_root) 117 } 118 119 fn load_config(workspace_root: &Path) -> Result<Config, String> { 120 let path = safe_workspace_file(workspace_root, CONFIG_PATH, false, "protocol codegen input")?; 121 let source = fs::read_to_string(&path) 122 .map_err(|error| format!("failed to read `{CONFIG_PATH}`: {error}"))?; 123 toml::from_str(&source).map_err(|error| format!("invalid `{CONFIG_PATH}`: {error}")) 124 } 125 126 fn validate_config(config: &Config) -> Result<(), String> { 127 if config.schema_version != 1 { 128 return Err("protocol codegen schema_version must be 1".to_owned()); 129 } 130 if config.package != PACKAGE { 131 return Err(format!("protocol codegen package must be `{PACKAGE}`")); 132 } 133 if config.source_hash_algorithm != HASH_ALGORITHM { 134 return Err(format!( 135 "protocol codegen source_hash_algorithm must be `{HASH_ALGORITHM}`" 136 )); 137 } 138 if config.inventory_path != "contracts/codegen/protocol_v1.inventory.json" 139 || config.inventory_sha256_path != "contracts/codegen/protocol_v1.inventory.sha256" 140 || config.packaged_inventory_path 141 != "crates/protocol/tests/fixtures/protocol_v1.inventory.json" 142 { 143 return Err("protocol codegen output paths are fixed by contract".to_owned()); 144 } 145 146 let mut actual = config 147 .source 148 .iter() 149 .map(|source| (source.module.as_str(), source.path.as_str())) 150 .collect::<Vec<_>>(); 151 actual.sort_unstable(); 152 if actual != EXPECTED_SOURCES { 153 return Err(format!( 154 "protocol codegen source inventory drifted: expected {EXPECTED_SOURCES:?}, found {actual:?}" 155 )); 156 } 157 if config 158 .source 159 .iter() 160 .map(|source| &source.module) 161 .collect::<BTreeSet<_>>() 162 .len() 163 != config.source.len() 164 { 165 return Err("protocol codegen modules must be unique".to_owned()); 166 } 167 168 let mut macro_generated = config 169 .macro_generated_type 170 .iter() 171 .map(|item| { 172 ( 173 item.module.as_str(), 174 item.path.as_str(), 175 item.macro_name.as_str(), 176 item.rust_name.as_str(), 177 item.kind.as_str(), 178 ) 179 }) 180 .collect::<Vec<_>>(); 181 macro_generated.sort_unstable(); 182 if macro_generated != EXPECTED_MACRO_GENERATED_TYPES { 183 return Err(format!( 184 "protocol macro-generated type inventory drifted: expected {EXPECTED_MACRO_GENERATED_TYPES:?}, found {macro_generated:?}" 185 )); 186 } 187 Ok(()) 188 } 189 190 fn protocol_schemas() -> Result<Vec<SchemaInventory>, String> { 191 let registry = radroots_protocol::schema::protocol_v1_registry() 192 .map_err(|error| format!("invalid protocol schema registry: {error}"))?; 193 Ok(registry 194 .descriptors() 195 .iter() 196 .map(|descriptor| SchemaInventory { 197 schema_id: descriptor.id().as_str().to_owned(), 198 module: descriptor.module().path().to_owned(), 199 generation: descriptor.module().generation(), 200 }) 201 .collect()) 202 } 203 204 fn render_outputs( 205 workspace_root: &Path, 206 config: &Config, 207 schemas: Vec<SchemaInventory>, 208 ) -> Result<Vec<GeneratedFile>, String> { 209 let mut sources = config.source.clone(); 210 sources.sort(); 211 let mut inventories = Vec::with_capacity(sources.len()); 212 for source in sources { 213 let path = safe_workspace_file(workspace_root, &source.path, false, "protocol DTO source")?; 214 let bytes = fs::read(&path).map_err(|error| { 215 format!( 216 "failed to read protocol DTO source `{}`: {error}", 217 source.path 218 ) 219 })?; 220 let text = std::str::from_utf8(&bytes).map_err(|error| { 221 format!( 222 "protocol DTO source `{}` is not UTF-8: {error}", 223 source.path 224 ) 225 })?; 226 let mut types = serialized_public_types(&source.module, &source.path, text)?; 227 for generated in config 228 .macro_generated_type 229 .iter() 230 .filter(|generated| generated.module == source.module && generated.path == source.path) 231 { 232 types.push(macro_generated_serialized_type( 233 generated, 234 source.path.as_str(), 235 text, 236 )?); 237 } 238 types.sort_by(|left, right| left.rust_path.cmp(&right.rust_path)); 239 if types 240 .windows(2) 241 .any(|pair| pair[0].rust_path == pair[1].rust_path) 242 { 243 return Err(format!( 244 "protocol DTO source `{}` exposes duplicate serialized type inventory entries", 245 source.path 246 )); 247 } 248 inventories.push(SourceInventory { 249 module: source.module.clone(), 250 path: source.path.clone(), 251 sha256: sha256_hex(&bytes), 252 types, 253 }); 254 } 255 256 let inventory = Inventory { 257 schema_version: 1, 258 generator: "radroots_xtask.protocol_codegen.v1", 259 package: PACKAGE, 260 source_hash_algorithm: HASH_ALGORITHM, 261 sources: inventories, 262 schemas, 263 }; 264 let mut inventory_bytes = serde_json::to_vec_pretty(&inventory) 265 .map_err(|error| format!("failed to serialize protocol DTO inventory: {error}"))?; 266 inventory_bytes.push(b'\n'); 267 let digest_bytes = format!("{}\n", sha256_hex(&inventory_bytes)).into_bytes(); 268 let packaged_inventory_bytes = inventory_bytes.clone(); 269 270 Ok(vec![ 271 GeneratedFile { 272 path: safe_workspace_file( 273 workspace_root, 274 &config.inventory_path, 275 true, 276 "protocol DTO inventory output", 277 )?, 278 display_path: config.inventory_path.clone(), 279 bytes: inventory_bytes, 280 }, 281 GeneratedFile { 282 path: safe_workspace_file( 283 workspace_root, 284 &config.inventory_sha256_path, 285 true, 286 "protocol DTO inventory digest output", 287 )?, 288 display_path: config.inventory_sha256_path.clone(), 289 bytes: digest_bytes, 290 }, 291 GeneratedFile { 292 path: safe_workspace_file( 293 workspace_root, 294 &config.packaged_inventory_path, 295 true, 296 "packaged protocol DTO inventory output", 297 )?, 298 display_path: config.packaged_inventory_path.clone(), 299 bytes: packaged_inventory_bytes, 300 }, 301 ]) 302 } 303 304 fn macro_generated_serialized_type( 305 config: &MacroGeneratedTypeConfig, 306 source_path: &str, 307 source: &str, 308 ) -> Result<TypeInventory, String> { 309 let syntax = syn::parse_file(source) 310 .map_err(|error| format!("failed to parse protocol DTO source `{source_path}`: {error}"))?; 311 let definition = syntax.items.iter().find_map(|item| match item { 312 Item::Macro(item) 313 if item 314 .ident 315 .as_ref() 316 .is_some_and(|name| name == &config.macro_name) => 317 { 318 Some(item.mac.tokens.to_string()) 319 } 320 _ => None, 321 }); 322 let Some(definition) = definition else { 323 return Err(format!( 324 "protocol DTO source `{source_path}` does not define configured macro `{}`", 325 config.macro_name 326 )); 327 }; 328 let invoked = syntax.items.iter().any(|item| { 329 matches!(item, Item::Macro(item) if item.ident.is_none() && item.mac.path.is_ident(&config.macro_name)) 330 }); 331 if !invoked { 332 return Err(format!( 333 "protocol DTO source `{source_path}` does not invoke configured macro `{}`", 334 config.macro_name 335 )); 336 } 337 let compact_definition = definition.split_whitespace().collect::<String>(); 338 for required in [ 339 format!("pub{}{}", config.kind, config.rust_name), 340 format!("implserde::Serializefor{}", config.rust_name), 341 format!("impl<'de>serde::Deserialize<'de>for{}", config.rust_name), 342 ] { 343 if !compact_definition.contains(&required) { 344 return Err(format!( 345 "configured macro `{}` in `{source_path}` does not prove serialized public {} `{}`: missing `{required}`", 346 config.macro_name, config.kind, config.rust_name 347 )); 348 } 349 } 350 Ok(TypeInventory { 351 rust_path: format!("radroots_protocol::{}::{}", config.module, config.rust_name), 352 kind: match config.kind.as_str() { 353 "enum" => "enum", 354 "struct" => "struct", 355 other => { 356 return Err(format!( 357 "unsupported macro-generated protocol DTO kind `{other}`" 358 )); 359 } 360 }, 361 }) 362 } 363 364 fn serialized_public_types( 365 module: &str, 366 source_path: &str, 367 source: &str, 368 ) -> Result<Vec<TypeInventory>, String> { 369 let syntax = syn::parse_file(source) 370 .map_err(|error| format!("failed to parse protocol DTO source `{source_path}`: {error}"))?; 371 let manual_serializers = syntax 372 .items 373 .iter() 374 .filter_map(manual_serialize_target) 375 .collect::<BTreeSet<_>>(); 376 let mut types = BTreeMap::new(); 377 for item in &syntax.items { 378 let (name, kind, visibility, attributes) = match item { 379 Item::Enum(item) => (&item.ident, "enum", &item.vis, &item.attrs), 380 Item::Struct(item) => (&item.ident, "struct", &item.vis, &item.attrs), 381 _ => continue, 382 }; 383 if !matches!(visibility, Visibility::Public(_)) { 384 continue; 385 } 386 let name = name.to_string(); 387 let derives_serialize = attributes.iter().any(|attribute| { 388 attribute 389 .meta 390 .to_token_stream() 391 .to_string() 392 .contains("Serialize") 393 }); 394 if derives_serialize || manual_serializers.contains(&name) { 395 types.insert( 396 name.clone(), 397 TypeInventory { 398 rust_path: format!("radroots_protocol::{module}::{name}"), 399 kind, 400 }, 401 ); 402 } 403 } 404 if types.is_empty() { 405 return Err(format!( 406 "protocol DTO source `{source_path}` exposes no serialized public types" 407 )); 408 } 409 Ok(types.into_values().collect()) 410 } 411 412 fn manual_serialize_target(item: &Item) -> Option<String> { 413 let Item::Impl(item) = item else { 414 return None; 415 }; 416 let (_, trait_path, _) = item.trait_.as_ref()?; 417 if trait_path.segments.last()?.ident != "Serialize" { 418 return None; 419 } 420 let Type::Path(self_type) = item.self_ty.as_ref() else { 421 return None; 422 }; 423 Some(self_type.path.segments.last()?.ident.to_string()) 424 } 425 426 fn check_outputs(generated: &[GeneratedFile]) -> Result<(), String> { 427 let stale = generated 428 .iter() 429 .filter_map(|output| match fs::read(&output.path) { 430 Ok(actual) if actual == output.bytes => None, 431 Ok(_) => Some(format!("stale `{}`", output.display_path)), 432 Err(error) if error.kind() == ErrorKind::NotFound => { 433 Some(format!("missing `{}`", output.display_path)) 434 } 435 Err(error) => Some(format!("unreadable `{}`: {error}", output.display_path)), 436 }) 437 .collect::<Vec<_>>(); 438 if stale.is_empty() { 439 Ok(()) 440 } else { 441 Err(format!( 442 "generated protocol DTO inventory is not fresh:\n- {}\nrun `cargo xtask generate protocol --write`", 443 stale.join("\n- ") 444 )) 445 } 446 } 447 448 fn write_outputs(generated: &[GeneratedFile]) -> Result<(), String> { 449 let mut staged = Vec::new(); 450 for output in generated { 451 if fs::read(&output.path).is_ok_and(|actual| actual == output.bytes) { 452 continue; 453 } 454 let parent = output 455 .path 456 .parent() 457 .ok_or_else(|| format!("generated output has no parent: `{}`", output.display_path))?; 458 let mut temporary = NamedTempFile::new_in(parent) 459 .map_err(|error| format!("failed to stage `{}`: {error}", output.display_path))?; 460 temporary 461 .write_all(&output.bytes) 462 .map_err(|error| format!("failed to stage `{}`: {error}", output.display_path))?; 463 set_generated_permissions(temporary.path()).map_err(|error| { 464 format!( 465 "failed to set generated permissions for `{}`: {error}", 466 output.display_path 467 ) 468 })?; 469 temporary 470 .as_file() 471 .sync_all() 472 .map_err(|error| format!("failed to sync `{}`: {error}", output.display_path))?; 473 staged.push((output, temporary)); 474 } 475 for (output, temporary) in staged { 476 temporary.persist(&output.path).map_err(|error| { 477 format!( 478 "failed to commit `{}`: {}", 479 output.display_path, error.error 480 ) 481 })?; 482 } 483 check_outputs(generated) 484 } 485 486 #[cfg(unix)] 487 fn set_generated_permissions(path: &Path) -> std::io::Result<()> { 488 use std::os::unix::fs::PermissionsExt; 489 490 fs::set_permissions(path, fs::Permissions::from_mode(0o644)) 491 } 492 493 #[cfg(not(unix))] 494 fn set_generated_permissions(_path: &Path) -> std::io::Result<()> { 495 Ok(()) 496 } 497 498 fn safe_workspace_file( 499 workspace_root: &Path, 500 relative: &str, 501 allow_missing: bool, 502 role: &str, 503 ) -> Result<PathBuf, String> { 504 let path = Path::new(relative); 505 if relative.is_empty() 506 || relative.contains('\\') 507 || path.is_absolute() 508 || !path 509 .components() 510 .all(|component| matches!(component, Component::Normal(_))) 511 { 512 return Err(format!( 513 "{role} path must be normalized and workspace-relative: `{relative}`" 514 )); 515 } 516 let mut current = workspace_root.to_path_buf(); 517 let count = path.components().count(); 518 for (index, component) in path.components().enumerate() { 519 let Component::Normal(segment) = component else { 520 return Err(format!("{role} path is not normalized: `{relative}`")); 521 }; 522 current.push(segment); 523 match current.symlink_metadata() { 524 Ok(metadata) if metadata.file_type().is_symlink() => { 525 return Err(format!( 526 "{role} path contains a symlink: `{}`", 527 current.display() 528 )); 529 } 530 Ok(metadata) if index + 1 == count && !metadata.is_file() => { 531 return Err(format!("{role} must be a regular file: `{relative}`")); 532 } 533 Ok(_) => {} 534 Err(error) 535 if error.kind() == ErrorKind::NotFound && allow_missing && index + 1 == count => {} 536 Err(error) => { 537 return Err(format!("failed to inspect {role} `{relative}`: {error}")); 538 } 539 } 540 } 541 Ok(current) 542 } 543 544 fn sha256_hex(bytes: &[u8]) -> String { 545 hex::encode(Sha256::digest(bytes)) 546 } 547 548 #[cfg(test)] 549 mod tests { 550 use super::*; 551 use tempfile::TempDir; 552 553 #[test] 554 fn public_serialized_type_inventory_is_sorted_and_excludes_native_types() { 555 let types = serialized_public_types( 556 "demo::v1", 557 "demo.rs", 558 r#" 559 #[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))] 560 pub struct Zebra { pub value: String } 561 pub struct NativeError; 562 pub struct Alpha(String); 563 impl serde::Serialize for Alpha {} 564 struct PrivateWire; 565 "#, 566 ) 567 .expect("inventory"); 568 assert_eq!( 569 types 570 .iter() 571 .map(|item| item.rust_path.as_str()) 572 .collect::<Vec<_>>(), 573 [ 574 "radroots_protocol::demo::v1::Alpha", 575 "radroots_protocol::demo::v1::Zebra" 576 ] 577 ); 578 } 579 580 #[test] 581 fn generated_bytes_and_freshness_check_are_deterministic() { 582 let workspace = TempDir::new().expect("workspace"); 583 let root = workspace.path(); 584 fs::create_dir_all(root.join("src")).expect("source directory"); 585 fs::create_dir_all(root.join("out")).expect("output directory"); 586 fs::create_dir_all(root.join("packaged")).expect("packaged output directory"); 587 fs::write( 588 root.join("src/types.rs"), 589 "#[derive(serde::Serialize)] pub struct Demo { pub value: String }\n", 590 ) 591 .expect("source"); 592 let config = Config { 593 schema_version: 1, 594 package: PACKAGE.to_owned(), 595 source_hash_algorithm: HASH_ALGORITHM.to_owned(), 596 inventory_path: "out/inventory.json".to_owned(), 597 inventory_sha256_path: "out/inventory.sha256".to_owned(), 598 packaged_inventory_path: "packaged/inventory.json".to_owned(), 599 source: vec![SourceConfig { 600 module: "demo::v1".to_owned(), 601 path: "src/types.rs".to_owned(), 602 }], 603 macro_generated_type: Vec::new(), 604 }; 605 let schemas = vec![SchemaInventory { 606 schema_id: "demo.message.v1".to_owned(), 607 module: "demo::v1".to_owned(), 608 generation: 1, 609 }]; 610 let first = render_outputs(root, &config, schemas.clone()).expect("first render"); 611 let second = render_outputs(root, &config, schemas).expect("second render"); 612 assert_eq!(first[0].bytes, second[0].bytes); 613 assert_eq!(first[1].bytes, second[1].bytes); 614 assert_eq!(first[2].bytes, second[2].bytes); 615 assert_eq!(first[0].bytes, first[2].bytes); 616 write_outputs(&first).expect("write"); 617 check_outputs(&second).expect("fresh"); 618 fs::write(root.join("out/inventory.json"), "stale\n").expect("drift"); 619 let error = check_outputs(&second).expect_err("reject drift"); 620 assert!(error.contains("stale `out/inventory.json`")); 621 write_outputs(&second).expect("restore outputs"); 622 fs::write(root.join("packaged/inventory.json"), "stale\n").expect("packaged drift"); 623 let error = check_outputs(&second).expect_err("reject packaged drift"); 624 assert!(error.contains("stale `packaged/inventory.json`")); 625 } 626 627 #[test] 628 fn normalized_paths_reject_escape_and_symlinks() { 629 let workspace = TempDir::new().expect("workspace"); 630 for path in ["", "../escape", "/absolute", "a/./b", "a\\b"] { 631 assert!(safe_workspace_file(workspace.path(), path, true, "fixture").is_err()); 632 } 633 } 634 635 #[test] 636 fn macro_generated_serialized_type_requires_definition_invocation_and_serde() { 637 let config = MacroGeneratedTypeConfig { 638 module: "demo::v1".to_owned(), 639 path: "demo.rs".to_owned(), 640 macro_name: "generated_ids".to_owned(), 641 rust_name: "GeneratedId".to_owned(), 642 kind: "enum".to_owned(), 643 }; 644 let source = r#" 645 macro_rules! generated_ids { 646 () => { 647 pub enum GeneratedId { One } 648 impl serde::Serialize for GeneratedId {} 649 impl<'de> serde::Deserialize<'de> for GeneratedId {} 650 }; 651 } 652 generated_ids!(); 653 "#; 654 let item = macro_generated_serialized_type(&config, "demo.rs", source) 655 .expect("macro-generated serialized type"); 656 assert_eq!(item.rust_path, "radroots_protocol::demo::v1::GeneratedId"); 657 assert_eq!(item.kind, "enum"); 658 659 let missing_serde = source.replace("impl serde::Serialize for GeneratedId {}", ""); 660 let error = macro_generated_serialized_type(&config, "demo.rs", &missing_serde) 661 .expect_err("missing serializer must fail closed"); 662 assert!(error.contains("does not prove serialized public enum `GeneratedId`")); 663 } 664 }