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
13const EXTENSION_STATE_FILE: &str = "extension-state.yaml";
15
16#[derive(Debug, Clone, Default, serde::Serialize, serde::Deserialize)]
19struct ExtensionState {
20 #[serde(default, skip_serializing_if = "Option::is_none")]
22 previous_default_registry: Option<types::Registry>,
23 #[serde(default, skip_serializing_if = "Option::is_none")]
25 project_file_path: Option<PathBuf>,
26 #[serde(default, skip_serializing_if = "Option::is_none")]
28 installed_extension: Option<types::ExtensionInfo>
29}
30
31fn 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
45fn 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
55fn 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
67fn 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
89fn 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
98fn 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
106fn 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
115fn 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
163pub 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
365pub 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
575fn 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}