Skip to main content

zoi_plugins/
extension.rs

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