Skip to main content

zoi_plugins/
extension.rs

1use std::fs;
2use std::path::PathBuf;
3
4use anyhow::{Result, anyhow};
5use mlua::LuaSerdeExt;
6use zoi_core::{config, pgp, types};
7use zoi_hooks as hooks;
8use zoi_lua;
9use zoi_resolver::{local, resolve};
10
11use crate::PluginManager;
12
13/// Filename used to store the persistent state of an extension.
14const EXTENSION_STATE_FILE: &str = "extension-state.yaml";
15
16/// Represents the persistent state of an installed extension, allowing for
17/// rollbacks.
18#[derive(Debug, Clone, Default, serde::Serialize, serde::Deserialize)]
19struct ExtensionState {
20    /// The default registry configuration before the extension was added.
21    #[serde(default, skip_serializing_if = "Option::is_none")]
22    previous_default_registry: Option<types::Registry>,
23    /// The path to the project file created by the extension, if any.
24    #[serde(default, skip_serializing_if = "Option::is_none")]
25    project_file_path: Option<PathBuf>,
26    /// Metadata about the installed extension itself.
27    #[serde(default, skip_serializing_if = "Option::is_none")]
28    installed_extension: Option<types::ExtensionInfo>
29}
30
31/// Returns the path to the extension state file for a given manifest.
32fn get_extension_state_path(
33    manifest: &types::InstallManifest
34) -> Result<PathBuf> {
35    let version_dir = local::get_package_version_dir(
36        manifest.scope,
37        &manifest.registry_handle,
38        &manifest.repo,
39        &manifest.name,
40        &manifest.version
41    )?;
42    Ok(version_dir.join(EXTENSION_STATE_FILE))
43}
44
45/// Writes the extension state to disk.
46fn write_extension_state(
47    manifest: &types::InstallManifest,
48    extension_state: &ExtensionState
49) -> Result<()> {
50    let state_path = get_extension_state_path(manifest)?;
51    fs::write(state_path, serde_yaml::to_string(extension_state)?)?;
52    Ok(())
53}
54
55/// Reads the extension state from disk, if it exists.
56fn read_extension_state(
57    manifest: &types::InstallManifest
58) -> Result<Option<ExtensionState>> {
59    let state_path = get_extension_state_path(manifest)?;
60    if !state_path.exists() {
61        return Ok(None);
62    }
63    let content = fs::read_to_string(state_path)?;
64    Ok(Some(serde_yaml::from_str(&content)?))
65}
66
67/// Restores the default registry to its previous state.
68fn restore_default_registry(
69    saved_state: Option<&ExtensionState>,
70    added_registry_url: &str
71) -> Result<()> {
72    if let Some(saved_state) = saved_state {
73        return config::set_user_default_registry(
74            saved_state.previous_default_registry.clone()
75        );
76    }
77
78    let user_config = config::read_user_config()?;
79    let should_clear = user_config
80        .default_registry
81        .as_ref()
82        .is_some_and(|registry| registry.url == added_registry_url);
83    if should_clear {
84        config::set_user_default_registry(None)?;
85    }
86    Ok(())
87}
88
89/// Checks if the extension state contains any data that needs to be persisted.
90fn extension_state_requires_persistence(
91    extension_state: &ExtensionState
92) -> bool {
93    extension_state.previous_default_registry.is_some()
94        || extension_state.project_file_path.is_some()
95        || extension_state.installed_extension.is_some()
96}
97
98/// Returns the path to the project file, either from saved state or the
99/// default.
100fn get_project_file_path(saved_state: Option<&ExtensionState>) -> PathBuf {
101    saved_state
102        .and_then(|state| state.project_file_path.clone())
103        .unwrap_or_else(|| PathBuf::from("zoi.yaml"))
104}
105
106/// Extracts the repository name from a git URL.
107fn get_repo_name_from_url(url: &str) -> &str {
108    url.trim_end_matches('/')
109        .split('/')
110        .next_back()
111        .unwrap_or_default()
112        .trim_end_matches(".git")
113}
114
115/// Reverts a specific change made by an extension.
116fn revert_extension_change(
117    change: &types::ExtensionChange,
118    saved_state: Option<&ExtensionState>
119) -> Result<()> {
120    match change {
121        types::ExtensionChange::RepoGit { add } => {
122            let repo_name = get_repo_name_from_url(add);
123            if !repo_name.is_empty() {
124                config::remove_git_repo(repo_name)?;
125            }
126        }
127        types::ExtensionChange::RegistryRepo { add } => {
128            restore_default_registry(saved_state, add)?;
129        }
130        types::ExtensionChange::RegistryAdd { add } => {
131            config::remove_added_registry(add)?;
132        }
133        types::ExtensionChange::RepoAdd { add } => {
134            config::remove_repo(add)?;
135        }
136        types::ExtensionChange::Project { add: _ } => {
137            let project_file_path = get_project_file_path(saved_state);
138            if project_file_path.exists() {
139                fs::remove_file(project_file_path)?;
140            }
141        }
142        types::ExtensionChange::Pgp { name, key: _ } => {
143            pgp::remove_key_by_name(name)?;
144        }
145        types::ExtensionChange::Plugin { name, script: _ } => {
146            let plugin_dir = crate::get_plugin_dir()?;
147            let plugin_path = plugin_dir.join(format!("{name}.lua"));
148            if plugin_path.exists() {
149                fs::remove_file(plugin_path)?;
150            }
151        }
152        types::ExtensionChange::Hook { name, content: _ } => {
153            let hooks_dir = hooks::global::get_user_hooks_dir()?;
154            let hook_path = hooks_dir.join(format!("{name}.hook.yaml"));
155            if hook_path.exists() {
156                fs::remove_file(hook_path)?;
157            }
158        }
159    }
160    Ok(())
161}
162
163/// Installs an extension, applying its declared changes to the system and
164/// registry configuration.
165///
166/// # Errors
167///
168/// Returns an error if the extension package cannot be resolved, if it's not an
169/// extension package, or if applying any of the changes fails.
170pub fn add(
171    ext_name: &str,
172    yes: bool,
173    plugin_manager: Option<&PluginManager>
174) -> Result<()> {
175    println!("Adding extension: {ext_name}");
176
177    let (pkg, _, _, pkg_lua_path, registry_handle, repo_type, _) =
178        resolve::resolve_package_and_version(ext_name, None, false, yes)?;
179
180    if pkg.package_type != types::PackageType::Extension {
181        return Err(anyhow!("'{ext_name}' is not an extension package."));
182    }
183
184    let mut pkg_val = None;
185    if let Some(pm) = plugin_manager {
186        let v = pm
187            .lua
188            .to_value(&pkg)
189            .map_err(|e: mlua::Error| anyhow!(e.to_string()))?;
190        pm.trigger_hook("on_pre_extension_add", Some(&v.clone()))?;
191        pkg_val = Some(v);
192    }
193
194    let Some(extension_info) = pkg.extension else {
195        return Err(anyhow!(
196            "'{ext_name}' is an extension package but contains no extension \
197             data."
198        ));
199    };
200    if extension_info.extension_type != "zoi" {
201        return Err(anyhow!(
202            "Unsupported extension type: {}",
203            extension_info.extension_type
204        ));
205    }
206    let has_registry_repo_change =
207        extension_info.changes.iter().any(|change| {
208            matches!(change, types::ExtensionChange::RegistryRepo { .. })
209        });
210    let has_project_change = extension_info
211        .changes
212        .iter()
213        .any(|change| matches!(change, types::ExtensionChange::Project { .. }));
214    let previous_default_registry = if has_registry_repo_change {
215        config::read_user_config()?.default_registry.clone()
216    } else {
217        None
218    };
219    let project_file_path = if has_project_change {
220        Some(std::env::current_dir()?.join("zoi.yaml"))
221    } else {
222        None
223    };
224    let extension_state = ExtensionState {
225        previous_default_registry,
226        project_file_path,
227        installed_extension: Some(extension_info.clone())
228    };
229
230    let manifest = types::InstallManifest {
231        name: pkg.name.clone(),
232        version: pkg.version.clone().unwrap_or_default(),
233        epoch: pkg.epoch,
234        revision: pkg.revision.clone(),
235        sub_package: None,
236        repo: pkg.repo.clone(),
237        repo_type: repo_type.unwrap_or_else(|| "unofficial".to_string()),
238        registry_handle: registry_handle.unwrap_or_default(),
239        package_type: pkg.package_type,
240        description: pkg.description.clone(),
241        reason: types::InstallReason::Direct,
242        scope: pkg.scope,
243        bins: None,
244        conflicts: None,
245        replaces: None,
246        provides: None,
247        backup: None,
248        installed_dependencies: vec![],
249        dependencies_v2: None,
250        chosen_options: vec![],
251        chosen_optionals: vec![],
252        install_method: None,
253        platform: zoi_core::utils::get_platform().unwrap_or_default(),
254        service: None,
255        installed_files: vec![],
256        installed_size: pkg.installed_size,
257        sandbox: None,
258        completions: None
259    };
260    let mut wrote_manifest = false;
261    let mut applied_changes = Vec::new();
262    let add_result = (|| -> Result<()> {
263        if extension_state_requires_persistence(&extension_state) {
264            local::write_manifest(&manifest)?;
265            local::persist_package_source(&manifest, &pkg_lua_path)?;
266            wrote_manifest = true;
267            write_extension_state(&manifest, &extension_state)?;
268        }
269
270        println!("Applying extension changes...");
271        for change in &extension_info.changes {
272            match change {
273                types::ExtensionChange::RepoGit { add } => {
274                    println!("Adding git repository: {add}");
275                    config::clone_git_repo(add)?;
276                }
277                types::ExtensionChange::RegistryRepo { add } => {
278                    println!("Setting registry to: {add}");
279                    config::set_default_registry(add)?;
280                }
281                types::ExtensionChange::RegistryAdd { add } => {
282                    println!("Adding registry: {add}");
283                    config::add_added_registry(add)?;
284                }
285                types::ExtensionChange::RepoAdd { add } => {
286                    println!("Adding repository: {add}");
287                    config::add_repo(add)?;
288                }
289                types::ExtensionChange::Project { add } => {
290                    let project_file_path =
291                        get_project_file_path(Some(&extension_state));
292                    println!("Creating {}...", project_file_path.display());
293                    if project_file_path.exists() {
294                        return Err(anyhow!(
295                            "A 'zoi.yaml' file already exists at '{}'. Please \
296                             remove it first.",
297                            project_file_path.display()
298                        ));
299                    }
300                    fs::write(&project_file_path, add)?;
301                }
302                types::ExtensionChange::Pgp { name, key } => {
303                    println!("Adding PGP key: {name} from {key}");
304                    if key.starts_with("http") {
305                        pgp::add_key_from_url(key, name, false)?;
306                    } else {
307                        pgp::add_key_from_fingerprint(key, name, false)?;
308                    }
309                }
310                types::ExtensionChange::Plugin { name, script } => {
311                    println!("Adding plugin: {name}");
312                    let plugin_dir = crate::get_plugin_dir()?;
313                    let plugin_path = plugin_dir.join(format!("{name}.lua"));
314                    fs::write(plugin_path, script)?;
315                }
316                types::ExtensionChange::Hook { name, content } => {
317                    println!("Adding global hook: {name}");
318                    let hooks_dir = hooks::global::get_user_hooks_dir()?;
319                    let hook_path = hooks_dir.join(format!("{name}.hook.yaml"));
320                    fs::write(hook_path, content)?;
321                }
322            }
323            applied_changes.push(change.clone());
324        }
325        if !wrote_manifest {
326            local::write_manifest(&manifest)?;
327            local::persist_package_source(&manifest, &pkg_lua_path)?;
328            wrote_manifest = true;
329        }
330        Ok(())
331    })();
332    if let Err(error) = add_result {
333        for change in applied_changes.iter().rev() {
334            if let Err(rollback_error) =
335                revert_extension_change(change, Some(&extension_state))
336            {
337                eprintln!(
338                    "Warning: failed to roll back extension change \
339                     {change:?}: {rollback_error}"
340                );
341            }
342        }
343        if wrote_manifest
344            && let Ok(package_dir) = local::get_package_dir(
345                manifest.scope,
346                &manifest.registry_handle,
347                &manifest.repo,
348                &manifest.name
349            )
350        {
351            let _ = fs::remove_dir_all(package_dir);
352        }
353        return Err(error);
354    }
355
356    if let (Some(pm), Some(v)) = (plugin_manager, pkg_val) {
357        pm.trigger_hook_nonfatal("on_post_extension_add", Some(&v));
358    }
359
360    println!("Successfully added extension '{ext_name}'.");
361
362    Ok(())
363}
364
365/// Uninstalls an extension and reverts all changes it made to the system and
366/// configuration.
367///
368/// # Errors
369///
370/// Returns an error if the extension is not installed, if it's not an extension
371/// package, or if reverting any of the changes fails.
372pub fn remove(
373    ext_name: &str,
374    yes: bool,
375    plugin_manager: Option<&PluginManager>
376) -> Result<()> {
377    println!("Removing extension: {ext_name}");
378
379    let request = resolve::parse_source_string(ext_name)?;
380    let mut candidates = Vec::new();
381    for scope in [
382        types::Scope::Project,
383        types::Scope::User,
384        types::Scope::System
385    ] {
386        candidates
387            .extend(local::find_installed_manifests_matching(&request, scope)?);
388    }
389
390    if candidates.is_empty() {
391        return Err(anyhow!("Extension '{ext_name}' is not installed."));
392    }
393
394    let manifest = select_candidate(ext_name, candidates, yes)?;
395    let scope = manifest.scope;
396
397    let mut manifest_val = None;
398    if let Some(pm) = plugin_manager {
399        let v = pm
400            .lua
401            .to_value(&manifest)
402            .map_err(|e: mlua::Error| anyhow!(e.to_string()))?;
403        pm.trigger_hook("on_pre_extension_remove", Some(&v.clone()))?;
404        manifest_val = Some(v);
405    }
406
407    if manifest.package_type != types::PackageType::Extension {
408        return Err(anyhow!("'{ext_name}' is not an extension package."));
409    }
410
411    let installed_source_path = local::get_package_source_path(&manifest)?;
412    let pkg = if installed_source_path.exists() {
413        let path = installed_source_path.to_str().ok_or_else(|| {
414            anyhow!("Stored package source path contains invalid UTF-8")
415        })?;
416        zoi_lua::parser::parse_lua_package(
417            path,
418            Some(&manifest.version),
419            Some(manifest.scope),
420            true
421        )?
422    } else {
423        let source = local::installed_manifest_source(&manifest);
424        let (pkg, _, _, _, _, _, _) = resolve::resolve_package_and_version(
425            &source,
426            Some(manifest.scope),
427            true,
428            yes
429        )?;
430        pkg
431    };
432
433    let extension_state = read_extension_state(&manifest)?;
434    let extension_info = extension_state
435        .as_ref()
436        .and_then(|state| state.installed_extension.clone())
437        .or(pkg.extension);
438
439    if let Some(extension_info) = extension_info {
440        if extension_info.extension_type != "zoi" {
441            return Err(anyhow!(
442                "Unsupported extension type: {}",
443                extension_info.extension_type
444            ));
445        }
446
447        println!("Reverting extension changes...");
448        for change in extension_info.changes.iter().rev() {
449            match change {
450                types::ExtensionChange::RepoGit { add } => {
451                    let repo_name = get_repo_name_from_url(add);
452                    if !repo_name.is_empty() {
453                        println!("Removing git repository: {repo_name}");
454                        if let Err(e) = revert_extension_change(
455                            change,
456                            extension_state.as_ref()
457                        ) {
458                            eprintln!(
459                                "Warning: failed to remove git repo \
460                                 '{repo_name}': {e}"
461                            );
462                        }
463                    }
464                }
465                types::ExtensionChange::RegistryRepo { add: _ } => {
466                    println!("Restoring previous default registry");
467                    if let Err(e) = revert_extension_change(
468                        change,
469                        extension_state.as_ref()
470                    ) {
471                        eprintln!(
472                            "Warning: failed to restore default registry: {e}"
473                        );
474                    }
475                }
476                types::ExtensionChange::RegistryAdd { add } => {
477                    println!("Removing registry: {add}");
478                    if let Err(e) = revert_extension_change(
479                        change,
480                        extension_state.as_ref()
481                    ) {
482                        eprintln!(
483                            "Warning: failed to remove registry '{add}': {e}"
484                        );
485                    }
486                }
487                types::ExtensionChange::RepoAdd { add } => {
488                    println!("Removing repository: {add}");
489                    if let Err(e) = revert_extension_change(
490                        change,
491                        extension_state.as_ref()
492                    ) {
493                        eprintln!(
494                            "Warning: failed to remove repo '{add}': {e}"
495                        );
496                    }
497                }
498                types::ExtensionChange::Project { add: _ } => {
499                    let project_file_path =
500                        get_project_file_path(extension_state.as_ref());
501                    println!("Removing {}...", project_file_path.display());
502                    if let Err(e) = revert_extension_change(
503                        change,
504                        extension_state.as_ref()
505                    ) {
506                        eprintln!(
507                            "Warning: failed to remove '{}': {e}",
508                            project_file_path.display()
509                        );
510                    }
511                }
512                types::ExtensionChange::Pgp { name, key: _ } => {
513                    println!("Removing PGP key: {name}");
514                    if let Err(e) = revert_extension_change(
515                        change,
516                        extension_state.as_ref()
517                    ) {
518                        eprintln!(
519                            "Warning: failed to remove PGP key '{name}': {e}"
520                        );
521                    }
522                }
523                types::ExtensionChange::Plugin { name, script: _ } => {
524                    println!("Removing plugin: {name}");
525                    if let Err(e) = revert_extension_change(
526                        change,
527                        extension_state.as_ref()
528                    ) {
529                        eprintln!(
530                            "Warning: failed to remove plugin '{name}': {e}"
531                        );
532                    }
533                }
534                types::ExtensionChange::Hook { name, content: _ } => {
535                    println!("Removing global hook: {name}");
536                    if let Err(e) = revert_extension_change(
537                        change,
538                        extension_state.as_ref()
539                    ) {
540                        eprintln!(
541                            "Warning: failed to remove global hook '{name}': \
542                             {e}"
543                        );
544                    }
545                }
546            }
547        }
548    } else {
549        return Err(anyhow!(
550            "'{ext_name}' is an extension package but contains no extension \
551             data."
552        ));
553    }
554
555    let package_dir = local::get_package_dir(
556        scope,
557        &manifest.registry_handle,
558        &manifest.repo,
559        &manifest.name
560    )?;
561
562    if package_dir.exists() {
563        fs::remove_dir_all(&package_dir)?;
564    }
565
566    if let (Some(pm), Some(v)) = (plugin_manager, manifest_val) {
567        pm.trigger_hook_nonfatal("on_post_extension_remove", Some(&v));
568    }
569
570    println!("Successfully removed extension '{ext_name}'.");
571
572    Ok(())
573}
574
575/// Prompts the user to select one extension if multiple candidates match the
576/// request.
577fn select_candidate(
578    package_name: &str,
579    candidates: Vec<types::InstallManifest>,
580    yes: bool
581) -> Result<types::InstallManifest> {
582    use colored::Colorize;
583    use comfy_table::Table;
584    use comfy_table::presets::UTF8_FULL;
585    use dialoguer::Select;
586    use dialoguer::theme::ColorfulTheme;
587
588    if candidates.is_empty() {
589        return Err(anyhow!("Package '{package_name}' is not installed."));
590    }
591    if candidates.len() == 1 {
592        return Ok(candidates
593            .into_iter()
594            .next()
595            .expect("Already checked length"));
596    }
597    if yes {
598        return Err(anyhow!(
599            "Package '{package_name}' matches multiple installed packages. \
600             Use an explicit source like '#handle@repo/name[:sub]@version'."
601        ));
602    }
603
604    let displays: Vec<_> = candidates
605        .iter()
606        .map(|m| {
607            let source = local::installed_manifest_source(m);
608            let scope_label = match m.scope {
609                types::Scope::User => "user",
610                types::Scope::System => "system",
611                types::Scope::Project => "project"
612            };
613            format!("{} ({}, v{})", source, scope_label, m.version)
614        })
615        .collect();
616
617    let mut table = Table::new();
618    table.load_style(UTF8_FULL);
619    table.set_header(vec!["#", "Source", "Version"]);
620    for (i, m) in candidates.iter().enumerate() {
621        table.add_row(vec![
622            (i + 1).to_string(),
623            local::installed_manifest_source(m),
624            m.version.clone(),
625        ]);
626    }
627    println!(
628        "Found multiple installed packages matching '{}'. Please choose one:",
629        package_name.cyan()
630    );
631    println!("{table}");
632
633    let selection = Select::with_theme(&ColorfulTheme::default())
634        .with_prompt("Select an installed package")
635        .items(&displays)
636        .default(0)
637        .interact()?;
638
639    Ok(candidates
640        .get(selection)
641        .cloned()
642        .expect("Invalid selection"))
643}