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}