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        epoch: pkg.epoch,
201        revision: pkg.revision.clone(),
202        sub_package: None,
203        repo: pkg.repo.clone(),
204        repo_type: repo_type.unwrap_or_else(|| "unofficial".to_string()),
205        registry_handle: registry_handle.unwrap_or_default(),
206        package_type: pkg.package_type,
207        description: pkg.description.clone(),
208        reason: types::InstallReason::Direct,
209        scope: pkg.scope,
210        bins: None,
211        conflicts: None,
212        replaces: None,
213        provides: None,
214        backup: None,
215        installed_dependencies: vec![],
216        dependencies_v2: None,
217        chosen_options: vec![],
218        chosen_optionals: vec![],
219        install_method: None,
220        platform: zoi_core::utils::get_platform().unwrap_or_default(),
221        service: None,
222        installed_files: vec![],
223        installed_size: pkg.installed_size,
224        sandbox: None,
225        completions: None,
226    };
227    let mut wrote_manifest = false;
228    let mut applied_changes = Vec::new();
229    let add_result = (|| -> Result<()> {
230        if extension_state_requires_persistence(&extension_state) {
231            local::write_manifest(&manifest)?;
232            local::persist_package_source(&manifest, &pkg_lua_path)?;
233            wrote_manifest = true;
234            write_extension_state(&manifest, &extension_state)?;
235        }
236
237        println!("Applying extension changes...");
238        for change in &extension_info.changes {
239            match change {
240                types::ExtensionChange::RepoGit { add } => {
241                    println!("Adding git repository: {}", add);
242                    config::clone_git_repo(add)?;
243                }
244                types::ExtensionChange::RegistryRepo { add } => {
245                    println!("Setting registry to: {}", add);
246                    config::set_default_registry(add)?;
247                }
248                types::ExtensionChange::RegistryAdd { add } => {
249                    println!("Adding registry: {}", add);
250                    config::add_added_registry(add)?;
251                }
252                types::ExtensionChange::RepoAdd { add } => {
253                    println!("Adding repository: {}", add);
254                    config::add_repo(add)?;
255                }
256                types::ExtensionChange::Project { add } => {
257                    let project_file_path = get_project_file_path(Some(&extension_state));
258                    println!("Creating {}...", project_file_path.display());
259                    if project_file_path.exists() {
260                        return Err(anyhow!(
261                            "A 'zoi.yaml' file already exists at '{}'. Please remove it first.",
262                            project_file_path.display()
263                        ));
264                    }
265                    fs::write(&project_file_path, add)?;
266                }
267                types::ExtensionChange::Pgp { name, key } => {
268                    println!("Adding PGP key: {} from {}", name, key);
269                    if key.starts_with("http") {
270                        pgp::add_key_from_url(key, name, false)?;
271                    } else {
272                        pgp::add_key_from_fingerprint(key, name, false)?;
273                    }
274                }
275                types::ExtensionChange::Plugin { name, script } => {
276                    println!("Adding plugin: {}", name);
277                    let plugin_dir = crate::get_plugin_dir()?;
278                    let plugin_path = plugin_dir.join(format!("{}.lua", name));
279                    fs::write(plugin_path, script)?;
280                }
281                types::ExtensionChange::Hook { name, content } => {
282                    println!("Adding global hook: {}", name);
283                    let hooks_dir = hooks::global::get_user_hooks_dir()?;
284                    let hook_path = hooks_dir.join(format!("{}.hook.yaml", name));
285                    fs::write(hook_path, content)?;
286                }
287            }
288            applied_changes.push(change.clone());
289        }
290        if !wrote_manifest {
291            local::write_manifest(&manifest)?;
292            local::persist_package_source(&manifest, &pkg_lua_path)?;
293            wrote_manifest = true;
294        }
295        Ok(())
296    })();
297    if let Err(error) = add_result {
298        for change in applied_changes.iter().rev() {
299            if let Err(rollback_error) = revert_extension_change(change, Some(&extension_state)) {
300                eprintln!(
301                    "Warning: failed to roll back extension change {:?}: {}",
302                    change, rollback_error
303                );
304            }
305        }
306        if wrote_manifest
307            && let Ok(package_dir) = local::get_package_dir(
308                manifest.scope,
309                &manifest.registry_handle,
310                &manifest.repo,
311                &manifest.name,
312            )
313        {
314            let _ = fs::remove_dir_all(package_dir);
315        }
316        return Err(error);
317    }
318
319    if let (Some(pm), Some(v)) = (plugin_manager, pkg_val) {
320        pm.trigger_hook_nonfatal("on_post_extension_add", Some(v));
321    }
322
323    println!("Successfully added extension '{}'.", ext_name);
324
325    Ok(())
326}
327
328pub fn remove(ext_name: &str, yes: bool, plugin_manager: Option<&PluginManager>) -> Result<()> {
329    println!("Removing extension: {}", ext_name);
330
331    let request = resolve::parse_source_string(ext_name)?;
332    let mut candidates = Vec::new();
333    for scope in [
334        types::Scope::Project,
335        types::Scope::User,
336        types::Scope::System,
337    ] {
338        candidates.extend(local::find_installed_manifests_matching(&request, scope)?);
339    }
340
341    if candidates.is_empty() {
342        return Err(anyhow!("Extension '{}' is not installed.", ext_name));
343    }
344
345    let manifest = select_candidate(ext_name, candidates, yes)?;
346    let scope = manifest.scope;
347
348    let mut manifest_val = None;
349    if let Some(pm) = plugin_manager {
350        let v = pm
351            .lua
352            .to_value(&manifest)
353            .map_err(|e: mlua::Error| anyhow!(e.to_string()))?;
354        pm.trigger_hook("on_pre_extension_remove", Some(v.clone()))?;
355        manifest_val = Some(v);
356    }
357
358    if manifest.package_type != types::PackageType::Extension {
359        return Err(anyhow!("'{}' is not an extension package.", ext_name));
360    }
361
362    let installed_source_path = local::get_package_source_path(&manifest)?;
363    let pkg = if installed_source_path.exists() {
364        let path = installed_source_path
365            .to_str()
366            .ok_or_else(|| anyhow!("Stored package source path contains invalid UTF-8"))?;
367        zoi_lua::parser::parse_lua_package(
368            path,
369            Some(&manifest.version),
370            Some(manifest.scope),
371            true,
372        )?
373    } else {
374        let source = local::installed_manifest_source(&manifest);
375        let (pkg, _, _, _, _, _, _) =
376            resolve::resolve_package_and_version(&source, Some(manifest.scope), true, yes)?;
377        pkg
378    };
379
380    let extension_state = read_extension_state(&manifest)?;
381    let extension_info = extension_state
382        .as_ref()
383        .and_then(|state| state.installed_extension.clone())
384        .or(pkg.extension);
385
386    if let Some(extension_info) = extension_info {
387        if extension_info.extension_type != "zoi" {
388            return Err(anyhow!(
389                "Unsupported extension type: {}",
390                extension_info.extension_type
391            ));
392        }
393
394        println!("Reverting extension changes...");
395        for change in extension_info.changes.iter().rev() {
396            match change {
397                types::ExtensionChange::RepoGit { add } => {
398                    let repo_name = get_repo_name_from_url(add);
399                    if !repo_name.is_empty() {
400                        println!("Removing git repository: {}", repo_name);
401                        if let Err(e) = revert_extension_change(change, extension_state.as_ref()) {
402                            eprintln!("Warning: failed to remove git repo '{}': {}", repo_name, e);
403                        }
404                    }
405                }
406                types::ExtensionChange::RegistryRepo { add: _ } => {
407                    println!("Restoring previous default registry");
408                    if let Err(e) = revert_extension_change(change, extension_state.as_ref()) {
409                        eprintln!("Warning: failed to restore default registry: {}", e);
410                    }
411                }
412                types::ExtensionChange::RegistryAdd { add } => {
413                    println!("Removing registry: {}", add);
414                    if let Err(e) = revert_extension_change(change, extension_state.as_ref()) {
415                        eprintln!("Warning: failed to remove registry '{}': {}", add, e);
416                    }
417                }
418                types::ExtensionChange::RepoAdd { add } => {
419                    println!("Removing repository: {}", add);
420                    if let Err(e) = revert_extension_change(change, extension_state.as_ref()) {
421                        eprintln!("Warning: failed to remove repo '{}': {}", add, e);
422                    }
423                }
424                types::ExtensionChange::Project { add: _ } => {
425                    let project_file_path = get_project_file_path(extension_state.as_ref());
426                    println!("Removing {}...", project_file_path.display());
427                    if let Err(e) = revert_extension_change(change, extension_state.as_ref()) {
428                        eprintln!(
429                            "Warning: failed to remove '{}': {}",
430                            project_file_path.display(),
431                            e
432                        );
433                    }
434                }
435                types::ExtensionChange::Pgp { name, key: _ } => {
436                    println!("Removing PGP key: {}", name);
437                    if let Err(e) = revert_extension_change(change, extension_state.as_ref()) {
438                        eprintln!("Warning: failed to remove PGP key '{}': {}", name, e);
439                    }
440                }
441                types::ExtensionChange::Plugin { name, script: _ } => {
442                    println!("Removing plugin: {}", name);
443                    if let Err(e) = revert_extension_change(change, extension_state.as_ref()) {
444                        eprintln!("Warning: failed to remove plugin '{}': {}", name, e);
445                    }
446                }
447                types::ExtensionChange::Hook { name, content: _ } => {
448                    println!("Removing global hook: {}", name);
449                    if let Err(e) = revert_extension_change(change, extension_state.as_ref()) {
450                        eprintln!("Warning: failed to remove global hook '{}': {}", name, e);
451                    }
452                }
453            }
454        }
455    } else {
456        return Err(anyhow!(
457            "'{}' is an extension package but contains no extension data.",
458            ext_name
459        ));
460    }
461
462    let package_dir = local::get_package_dir(
463        scope,
464        &manifest.registry_handle,
465        &manifest.repo,
466        &manifest.name,
467    )?;
468
469    if package_dir.exists() {
470        fs::remove_dir_all(&package_dir)?;
471    }
472
473    if let (Some(pm), Some(v)) = (plugin_manager, manifest_val) {
474        pm.trigger_hook_nonfatal("on_post_extension_remove", Some(v));
475    }
476
477    println!("Successfully removed extension '{}'.", ext_name);
478
479    Ok(())
480}
481
482fn select_candidate(
483    package_name: &str,
484    candidates: Vec<types::InstallManifest>,
485    yes: bool,
486) -> Result<types::InstallManifest> {
487    if candidates.is_empty() {
488        return Err(anyhow!("Package '{}' is not installed.", package_name));
489    }
490    if candidates.len() == 1 {
491        return Ok(candidates.into_iter().next().unwrap());
492    }
493    if yes {
494        return Err(anyhow!(
495            "Package '{}' matches multiple installed packages. Use an explicit source like '#handle@repo/name[:sub]@version'.",
496            package_name
497        ));
498    }
499
500    use colored::*;
501    use comfy_table::{Table, presets::UTF8_FULL};
502    use dialoguer::{Select, theme::ColorfulTheme};
503
504    let displays: Vec<_> = candidates
505        .iter()
506        .map(|m| {
507            let source = local::installed_manifest_source(m);
508            let scope_label = match m.scope {
509                types::Scope::User => "user",
510                types::Scope::System => "system",
511                types::Scope::Project => "project",
512            };
513            format!("{} ({}, v{})", source, scope_label, m.version)
514        })
515        .collect();
516
517    let mut table = Table::new();
518    table.load_preset(UTF8_FULL);
519    table.set_header(vec!["#", "Source", "Version"]);
520    for (i, m) in candidates.iter().enumerate() {
521        table.add_row(vec![
522            (i + 1).to_string(),
523            local::installed_manifest_source(m),
524            m.version.clone(),
525        ]);
526    }
527    println!(
528        "Found multiple installed packages matching '{}'. Please choose one:",
529        package_name.cyan()
530    );
531    println!("{table}");
532
533    let selection = Select::with_theme(&ColorfulTheme::default())
534        .with_prompt("Select an installed package")
535        .items(&displays)
536        .default(0)
537        .interact()?;
538
539    Ok(candidates[selection].clone())
540}