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}