Skip to main content

anza_xtask/commands/
publish.rs

1use {
2    crate::utils::{check_docker_available, get_git_root_path},
3    anyhow::{anyhow, Result},
4    cargo_metadata::{MetadataCommand, PackageId},
5    clap::{Args, Subcommand},
6    log::info,
7    scopeguard::defer,
8    std::{
9        collections::{HashMap, HashSet},
10        fs,
11        path::{Path, PathBuf},
12        process::Command,
13        sync::{Arc, RwLock},
14        thread,
15    },
16    toml_edit::{value, DocumentMut},
17};
18
19#[derive(Debug, Clone, serde::Serialize)]
20pub struct PackageInfo {
21    pub name: String,
22    pub path: std::path::PathBuf,
23    pub dependencies: HashSet<PackageId>,
24}
25
26pub struct PublishOrderData {
27    pub levels: Vec<Vec<PackageId>>,
28    pub id_to_level: std::collections::HashMap<PackageId, usize>,
29    pub id_to_package_info: HashMap<PackageId, PackageInfo>,
30}
31
32#[derive(Debug, Clone, clap::ValueEnum)]
33pub enum OutputFormat {
34    Json,
35    Tree,
36}
37
38#[derive(Subcommand)]
39pub enum PublishSubcommand {
40    #[command(about = "Print the publish order")]
41    Order {
42        #[arg(long, value_enum, default_value = "json")]
43        format: OutputFormat,
44    },
45    #[command(about = "Test the publish process")]
46    Test,
47}
48
49#[derive(Args)]
50pub struct CommandArgs {
51    #[arg(long, default_value = "Cargo.toml")]
52    pub manifest_path: String,
53
54    #[command(subcommand)]
55    pub subcommand: PublishSubcommand,
56}
57
58pub fn run(args: CommandArgs) -> Result<()> {
59    match args.subcommand {
60        PublishSubcommand::Order { format } => match format {
61            OutputFormat::Json => publish_order_json(&args.manifest_path)?,
62            OutputFormat::Tree => publish_order_tree(&args.manifest_path)?,
63        },
64        PublishSubcommand::Test => {
65            publish_test(&args.manifest_path)?;
66        }
67    }
68    Ok(())
69}
70
71pub fn compute_publish_order_data(manifest_path: &str) -> Result<PublishOrderData> {
72    let mut cmd = MetadataCommand::new();
73    cmd.features(cargo_metadata::CargoOpt::AllFeatures);
74    cmd.manifest_path(manifest_path);
75    let metadata = cmd.exec()?;
76
77    let workspace_member_ids: HashSet<&PackageId> = metadata.workspace_members.iter().collect();
78
79    let mut id_to_package_info: HashMap<PackageId, PackageInfo> = HashMap::new();
80    for pkg in metadata.packages.iter() {
81        // skip packages that are not part of the workspace
82        if !workspace_member_ids.contains(&pkg.id) {
83            continue;
84        }
85
86        // skip packages that no need to be published
87        if let Some(registries) = &pkg.publish {
88            if registries.is_empty() {
89                continue;
90            }
91        }
92
93        let path = Path::new(&pkg.manifest_path)
94            .parent()
95            .map(|p| p.to_path_buf())
96            .unwrap_or_else(|| PathBuf::from("."));
97
98        id_to_package_info.insert(
99            pkg.id.clone(),
100            PackageInfo {
101                name: pkg.name.clone().to_string(),
102                path,
103                dependencies: HashSet::new(),
104            },
105        );
106    }
107
108    // build dependency relationships
109    if let Some(resolve) = &metadata.resolve {
110        for node in &resolve.nodes {
111            // only process packages that are in our workspace
112            if let Some(mut package_info) = id_to_package_info.get(&node.id).cloned() {
113                for dep in node.deps.iter() {
114                    // skip self dependencies
115                    if dep.pkg == node.id {
116                        continue;
117                    }
118                    // skip dev-only dependencies - they don't affect publish order
119                    if dep
120                        .dep_kinds
121                        .iter()
122                        .all(|dk| dk.kind == cargo_metadata::DependencyKind::Development)
123                    {
124                        continue;
125                    }
126                    if id_to_package_info.contains_key(&dep.pkg) {
127                        package_info.dependencies.insert(dep.pkg.clone());
128                    }
129                }
130                id_to_package_info.insert(node.id.clone(), package_info);
131            }
132        }
133    }
134
135    let mut levels: Vec<Vec<PackageId>> = Vec::new();
136    let mut processed: HashSet<PackageId> = HashSet::new();
137    let mut id_to_level: HashMap<PackageId, usize> = HashMap::new();
138
139    loop {
140        let mut current_level = vec![];
141        // find all packages that have all their dependencies processed
142        for (package_id, package_info) in id_to_package_info.iter() {
143            if processed.contains(package_id) {
144                continue;
145            }
146            if package_info
147                .dependencies
148                .iter()
149                .all(|dep| processed.contains(dep))
150            {
151                current_level.push(package_id.clone());
152            }
153        }
154
155        if current_level.is_empty() {
156            break;
157        }
158        current_level.sort();
159
160        // add the current level to the levels vector
161        for package_id in current_level.iter().cloned() {
162            id_to_level.insert(package_id, levels.len());
163        }
164
165        levels.push(current_level.to_vec());
166
167        // mark the packages in the current level as processed
168        for package_id in current_level.iter().cloned() {
169            processed.insert(package_id);
170        }
171    }
172
173    // check for unprocessed packages
174    let mut unprocessed_packages = vec![];
175    for package_id in id_to_package_info.keys() {
176        if !processed.contains(package_id) {
177            let package_info = id_to_package_info.get(package_id).unwrap();
178            unprocessed_packages.push(package_info.name.clone());
179        }
180    }
181    if !unprocessed_packages.is_empty() {
182        return Err(anyhow!(
183            "Unprocessed packages found: {unprocessed_packages:?}",
184        ));
185    }
186
187    Ok(PublishOrderData {
188        levels,
189        id_to_level,
190        id_to_package_info,
191    })
192}
193
194pub fn publish_order_json(manifest_path: &str) -> Result<()> {
195    let publish_order_data = compute_publish_order_data(manifest_path)?;
196
197    let mut output = vec![];
198    for level in publish_order_data.levels.iter() {
199        let mut level_output = vec![];
200        for package_id in level.iter() {
201            let package_info = publish_order_data
202                .id_to_package_info
203                .get(package_id)
204                .unwrap();
205            level_output.push(package_info.to_owned());
206        }
207        output.push(level_output);
208    }
209
210    let json = serde_json::to_string(&output)?;
211    println!("{json}");
212
213    Ok(())
214}
215
216pub fn publish_order_tree(manifest_path: &str) -> Result<()> {
217    let publish_order_data = compute_publish_order_data(manifest_path)?;
218
219    let total_packages = publish_order_data
220        .levels
221        .iter()
222        .map(|level| level.len())
223        .sum::<usize>();
224    let total_levels = publish_order_data.levels.len();
225
226    println!("๐Ÿ“ฆ Total packages: {total_packages}");
227    println!("๐ŸŒณ Total levels: {total_levels}");
228    println!();
229
230    for (level, package_ids) in publish_order_data.levels.iter().enumerate() {
231        println!(
232            "L{}: ({} package(s))",
233            level.saturating_add(1),
234            package_ids.len()
235        );
236
237        for package_id in package_ids {
238            let package_info = publish_order_data
239                .id_to_package_info
240                .get(package_id)
241                .unwrap();
242            let package_name = &package_info.name;
243            let dependencies = &package_info.dependencies;
244
245            println!("  {package_name}");
246
247            if !dependencies.is_empty() {
248                // build a map of level -> dependencies
249                let mut dependencies_by_level: HashMap<usize, Vec<String>> = HashMap::new();
250                for dependency_package_id in dependencies.iter() {
251                    if let Some(&dependency_level) =
252                        publish_order_data.id_to_level.get(dependency_package_id)
253                    {
254                        let dependency_package_name = &publish_order_data
255                            .id_to_package_info
256                            .get(dependency_package_id)
257                            .unwrap()
258                            .name;
259                        dependencies_by_level
260                            .entry(dependency_level)
261                            .or_default()
262                            .push(dependency_package_name.clone());
263                    }
264                }
265
266                // sort levels
267                let mut sorted_levels: Vec<_> = dependencies_by_level.keys().copied().collect();
268                sorted_levels.sort();
269
270                for dependency_level in sorted_levels {
271                    println!(
272                        "    L{}: {:?}",
273                        dependency_level.saturating_add(1),
274                        dependencies_by_level[&dependency_level]
275                    );
276                }
277            }
278        }
279        println!();
280    }
281
282    Ok(())
283}
284
285fn write_custom_registry_config() -> Result<()> {
286    let git_root = get_git_root_path()?;
287    let config_file_path = git_root.join(".cargo/config.toml");
288    let content = fs::read_to_string(&config_file_path)
289        .map_err(|e| anyhow!("Failed to read config file: {e}"))?;
290    let mut doc = content
291        .parse::<DocumentMut>()
292        .map_err(|e| anyhow!("Failed to parse config file: {e}"))?;
293
294    let mut credential_provider = toml_edit::Array::new();
295    credential_provider.push("cargo:token");
296
297    doc["registries"]["kellnr"]["index"] = value("sparse+http://127.0.0.1:8000/api/v1/crates/");
298    doc["registries"]["kellnr"]["credential-provider"] = value(credential_provider);
299    doc["registries"]["kellnr"]["token"] = value("Zy9HhJ02RJmg0GCrgLfaCVfU6IwDfhXD");
300
301    fs::write(&config_file_path, doc.to_string())
302        .map_err(|e| anyhow!("Failed to write config file: {e}"))?;
303    Ok(())
304}
305
306fn start_docker_registry() -> Result<String> {
307    let output = Command::new("docker")
308        .args([
309            "run",
310            "--rm",
311            "-d",
312            "--name",
313            "kellnr",
314            "-p",
315            "8000:8000",
316            "ghcr.io/kellnr/kellnr:5",
317        ])
318        .output()
319        .map_err(|e| anyhow!("Failed to start docker container: {e}"))?;
320
321    if !output.status.success() {
322        let stderr = String::from_utf8_lossy(&output.stderr);
323        return Err(anyhow!("Failed to start docker container: {stderr}"));
324    }
325
326    let container_id = String::from_utf8_lossy(&output.stdout).trim().to_string();
327    Ok(container_id)
328}
329
330fn publish_test(manifest_path: &str) -> Result<()> {
331    defer! {
332        let git_root = get_git_root_path().unwrap();
333        let config_file_path = git_root.join(".cargo/config.toml");
334        info!("๐Ÿงน Cleanup: git checkout {:?}", config_file_path.to_str().unwrap());
335        Command::new("git")
336            .args(["checkout", &config_file_path.to_string_lossy()])
337            .output()
338            .map_err(|e| anyhow::anyhow!("Failed to run git checkout: {e}")).unwrap();
339    }
340
341    info!("checking docker");
342    check_docker_available()?;
343
344    info!("writing custom registry config to config file");
345    write_custom_registry_config()?;
346
347    info!("starting self-hosted kellnr registry");
348    let container_id = start_docker_registry()?;
349    info!("kellnr registry started: {container_id}");
350    defer! {
351        info!("๐Ÿงน Cleanup: stopping self-hosted kellnr registry");
352        Command::new("docker")
353            .args(["stop", &container_id])
354            .output()
355            .map_err(|e| anyhow::anyhow!("Failed to stop docker container: {e}")).unwrap();
356        Command::new("docker")
357            .args(["rm", &container_id])
358            .output()
359            .map_err(|e| anyhow::anyhow!("Failed to remove docker container: {e}")).unwrap();
360    }
361
362    info!("starting publish process");
363    let publish_order_data = compute_publish_order_data(manifest_path)?;
364    info!("total levels: {}", publish_order_data.levels.len());
365    info!(
366        "total packages: {}",
367        publish_order_data
368            .levels
369            .iter()
370            .map(|level| level.len())
371            .sum::<usize>()
372    );
373    for (level, package_ids) in publish_order_data.levels.iter().enumerate() {
374        info!("publishing level: {}", level.saturating_add(1));
375        info!("publishing {} package(s)", package_ids.len());
376        let mut handles = vec![];
377        for package_id in package_ids.iter() {
378            let package_info = publish_order_data
379                .id_to_package_info
380                .get(package_id)
381                .unwrap();
382
383            let package_name = package_info.name.clone();
384            let package_path = package_info.path.clone();
385
386            info!("  publishing package: {package_name}");
387            let handle = thread::spawn(move || -> Result<String> {
388                publish_package_with_docker(package_name.clone(), &package_path)
389                    .map_err(|e| anyhow!("Failed to publish package {package_name}: {e}"))?;
390                info!("    โœ… {package_name} published");
391                Ok(package_name)
392            });
393            handles.push(handle);
394        }
395
396        // wait for all threads and check for errors
397        let mut errors = vec![];
398        let manifest_lock = Arc::new(RwLock::new(()));
399        for handle in handles {
400            match handle.join() {
401                Ok(result) => {
402                    if let Ok(package_name) = result {
403                        update_workspace_manifest_registry(
404                            manifest_path,
405                            &package_name,
406                            &manifest_lock,
407                        )?;
408                    } else if let Err(e) = result {
409                        errors.push(e);
410                    }
411                }
412                Err(panic_payload) => {
413                    errors.push(anyhow!("Thread panicked: {panic_payload:?}"));
414                }
415            }
416        }
417        if !errors.is_empty() {
418            return Err(anyhow!(
419                "Failed to publish {} package(s) in level {}:\n{}",
420                errors.len(),
421                level.saturating_add(1),
422                errors
423                    .iter()
424                    .map(|e| format!("  - {e}"))
425                    .collect::<Vec<_>>()
426                    .join("\n")
427            ));
428        }
429    }
430    Ok(())
431}
432
433fn publish_package_with_docker(package_name: String, package_path: &Path) -> Result<String> {
434    let git_root = get_git_root_path()?;
435    let relative_package_path = package_path.strip_prefix(&git_root).unwrap_or(package_path);
436    let manifest_path = relative_package_path.join("Cargo.toml");
437    let output = Command::new(git_root.join("ci/docker-run-default-image.sh"))
438        .args([
439            "cargo",
440            "publish",
441            "--manifest-path",
442            &manifest_path.to_string_lossy(),
443            "--registry",
444            "kellnr",
445            "--allow-dirty",
446        ])
447        .current_dir(&git_root)
448        .env("EXTRA_DOCKER_RUN_ARGS", "--network container:kellnr")
449        .output()
450        .map_err(|e| anyhow::anyhow!("Failed to publish package: {e}"))
451        .unwrap();
452    if !output.status.success() {
453        let stdout = String::from_utf8_lossy(&output.stdout);
454        let stderr = String::from_utf8_lossy(&output.stderr);
455        return Err(anyhow!("Failed to publish package: {stderr}\n{stdout}"));
456    }
457    Ok(package_name)
458}
459
460fn update_workspace_manifest_registry(
461    manifest_path: &str,
462    package_name: &str,
463    manifest_lock: &RwLock<()>,
464) -> Result<()> {
465    // get the write lock
466    let _lock = manifest_lock
467        .write()
468        .map_err(|e| anyhow::anyhow!("Failed to get write lock: {e}"))?;
469
470    // read the manifest file
471    let content = fs::read_to_string(manifest_path)
472        .map_err(|e| anyhow::anyhow!("Failed to read manifest at {manifest_path}: {e}"))?;
473
474    // parse the toml document
475    let mut doc = content
476        .parse::<DocumentMut>()
477        .map_err(|e| anyhow::anyhow!("Failed to parse TOML: {e}"))?;
478
479    // add kellnr registry to the package
480    doc["workspace"]["dependencies"][package_name]["registry"] = value("kellnr");
481
482    // write back to file
483    fs::write(manifest_path, doc.to_string())
484        .map_err(|e| anyhow::anyhow!("Failed to write manifest: {e}"))?;
485
486    Ok(())
487}
488
489#[cfg(test)]
490mod tests {
491    use super::*;
492
493    #[test]
494    fn test_dependencies_are_in_earlier_levels() {
495        let manifest = "tests/dummy-workspace-publish-test/Cargo.toml";
496        let result = compute_publish_order_data(manifest);
497        assert!(result.is_ok(), "Should successfully compute order");
498        let data = result.unwrap();
499
500        let mut pkg_to_level = std::collections::HashMap::new();
501        for (level_idx, level) in data.levels.iter().enumerate() {
502            for pkg_id in level {
503                if let Some(pkg_info) = data.id_to_package_info.get(pkg_id) {
504                    pkg_to_level.insert(&pkg_info.name, level_idx);
505                }
506            }
507        }
508
509        for (level_idx, level) in data.levels.iter().enumerate() {
510            for pkg_id in level {
511                if let Some(pkg_info) = data.id_to_package_info.get(pkg_id) {
512                    for dep_id in &pkg_info.dependencies {
513                        if let Some(dep_info) = data.id_to_package_info.get(dep_id) {
514                            if let Some(&dep_level) = pkg_to_level.get(&dep_info.name) {
515                                assert!(dep_level <= level_idx,
516                                    "Dependency {} (level {}) should be published before or with {} (level {})",
517                                    dep_info.name, dep_level, pkg_info.name, level_idx);
518                            }
519                        }
520                    }
521                }
522            }
523        }
524    }
525
526    #[test]
527    fn test_publish_order_json_output() {
528        let manifest = "tests/dummy-workspace/Cargo.toml";
529        let result = publish_order_json(manifest);
530        assert!(result.is_ok(), "JSON output should succeed");
531    }
532
533    #[test]
534    fn test_invalid_manifest_path() {
535        let result = compute_publish_order_data("nonexistent/Cargo.toml");
536        assert!(result.is_err(), "Should fail with invalid manifest path");
537    }
538
539    #[test]
540    fn test_run_with_json_format() {
541        let args = CommandArgs {
542            manifest_path: "tests/dummy-workspace/Cargo.toml".to_string(),
543            subcommand: PublishSubcommand::Order {
544                format: OutputFormat::Json,
545            },
546        };
547        let result = run(args);
548        assert!(result.is_ok(), "Should succeed with JSON format");
549    }
550
551    #[test]
552    fn test_run_with_tree_format() {
553        let args = CommandArgs {
554            manifest_path: "tests/dummy-workspace/Cargo.toml".to_string(),
555            subcommand: PublishSubcommand::Order {
556                format: OutputFormat::Tree,
557            },
558        };
559        let result = run(args);
560        assert!(result.is_ok(), "Should succeed with Tree format");
561    }
562
563    #[test]
564    fn test_publish_excluded_packages() {
565        let manifest = "tests/dummy-workspace-publish-excluded/Cargo.toml";
566        let result = compute_publish_order_data(manifest);
567        assert!(result.is_ok());
568        let data = result.unwrap();
569
570        let names: Vec<&str> = data
571            .id_to_package_info
572            .values()
573            .map(|p| p.name.as_str())
574            .collect();
575        assert!(
576            names.contains(&"publishable"),
577            "publishable should be included"
578        );
579        assert!(
580            !names.contains(&"excluded"),
581            "excluded (publish=[]) should be filtered out"
582        );
583        assert_eq!(data.levels.len(), 1);
584    }
585
586    #[test]
587    fn test_publish_order_exact_levels() {
588        // dummy-workspace-publish-test has: a (no deps), b (depends on a),
589        // c (depends on b), d (depends on a and c)
590        // expected levels: 0=[a], 1=[b], 2=[c], 3=[d]
591        let manifest = "tests/dummy-workspace-publish-test/Cargo.toml";
592        let result = compute_publish_order_data(manifest);
593        assert!(result.is_ok());
594        let data = result.unwrap();
595
596        assert_eq!(data.levels.len(), 4);
597
598        let mut name_to_level: std::collections::HashMap<String, usize> = Default::default();
599        for (level_idx, level) in data.levels.iter().enumerate() {
600            for pkg_id in level {
601                name_to_level.insert(data.id_to_package_info[pkg_id].name.clone(), level_idx);
602            }
603        }
604
605        assert_eq!(name_to_level["a"], 0);
606        assert_eq!(name_to_level["b"], 1);
607        assert_eq!(name_to_level["c"], 2);
608        assert_eq!(
609            name_to_level["d"], 3,
610            "d depends on a and c, so it must be last"
611        );
612    }
613
614    #[test]
615    fn test_publish_order_tree_with_dependencies() {
616        // uses a workspace with inter-dependencies to exercise the dependency display path
617        let manifest = "tests/dummy-workspace-publish-test/Cargo.toml";
618        let result = publish_order_tree(manifest);
619        assert!(
620            result.is_ok(),
621            "Tree output with dependencies should succeed"
622        );
623    }
624
625    #[test]
626    fn test_update_workspace_manifest_registry() {
627        use std::sync::RwLock;
628        use tempfile::NamedTempFile;
629
630        // Create a temporary manifest file with workspace dependencies
631        let manifest_content = r#"
632[workspace]
633dependencies = { test-package = { version = "1.0.0", path = "../test" } }
634"#;
635
636        let mut temp_file = NamedTempFile::new().unwrap();
637        std::io::Write::write_all(&mut temp_file, manifest_content.as_bytes()).unwrap();
638        let temp_path = temp_file.path().to_str().unwrap();
639
640        let lock = RwLock::new(());
641        let result = update_workspace_manifest_registry(temp_path, "test-package", &lock);
642
643        assert!(result.is_ok(), "Should successfully update manifest");
644
645        // Verify the registry was added
646        let updated_content = fs::read_to_string(temp_path).unwrap();
647        assert!(
648            updated_content.contains("registry"),
649            "Should contain registry field"
650        );
651        assert!(
652            updated_content.contains("kellnr"),
653            "Should set registry to kellnr"
654        );
655    }
656}