bake_agent_context/agent/context/
installer.rs1use bake::{Error, Result};
5use serde::Deserialize;
6use std::collections::{HashMap, HashSet};
7use std::env;
8use std::fs;
9use std::path::{Component, Path, PathBuf};
10use std::process::Command;
11
12#[derive(Clone, Debug)]
14pub struct ContextPackage {
15 pub name: String,
16 pub version: String,
17 pub description: Option<String>,
18 pub context_path: PathBuf,
19 selector: String,
20}
21
22impl ContextPackage {
23 pub fn selector(&self) -> &str {
25 &self.selector
26 }
27}
28
29#[derive(Clone, Debug, Eq, PartialEq)]
31pub struct ContextFile {
32 pub path: PathBuf,
33}
34
35#[derive(Clone, Debug)]
37pub struct Installer {
38 root: PathBuf,
39 context_path: PathBuf,
40 packages: Vec<ContextPackage>,
41}
42
43impl Installer {
44 pub fn new(root: impl Into<PathBuf>) -> Result<Self> {
46 let root = root.into();
47 let manifest = root.join("Cargo.toml");
48 let cargo = env::var_os("CARGO").unwrap_or_else(|| "cargo".into());
49 let mut command = Command::new(cargo);
50 command
51 .args([
52 "metadata",
53 "--format-version",
54 "1",
55 "--locked",
56 "--manifest-path",
57 ])
58 .arg(&manifest)
59 .current_dir(&root);
60
61 let output = command.output().map_err(|error| {
62 Error::new(format!(
63 "cannot run cargo metadata for {}: {error}",
64 manifest.display()
65 ))
66 })?;
67 if !output.status.success() {
68 let details = String::from_utf8_lossy(&output.stderr).trim().to_owned();
69 return Err(Error::new(format!(
70 "cargo metadata failed for {} ({}): {}",
71 manifest.display(),
72 output.status,
73 if details.is_empty() {
74 "run cargo check to resolve and lock the project's dependencies".to_owned()
75 } else {
76 details
77 }
78 )));
79 }
80
81 let metadata: CargoMetadata = serde_json::from_slice(&output.stdout)
82 .map_err(|error| Error::new(format!("cannot parse cargo metadata: {error}")))?;
83 let workspace_members: HashSet<_> = metadata.workspace_members.into_iter().collect();
84 let resolved_packages: HashSet<_> = metadata
85 .resolve
86 .map(|resolve| {
87 resolve
88 .nodes
89 .into_iter()
90 .map(|node| node.package_id)
91 .collect()
92 })
93 .unwrap_or_default();
94
95 let mut candidates = Vec::new();
96 for package in metadata.packages {
97 if workspace_members.contains(&package.package_id)
98 || (!resolved_packages.is_empty()
99 && !resolved_packages.contains(&package.package_id))
100 {
101 continue;
102 }
103
104 let Some(package_root) = package.manifest_path.parent() else {
105 continue;
106 };
107 let context_path = package_root.join("context");
108 let context_metadata = match fs::symlink_metadata(&context_path) {
109 Ok(metadata) if metadata.file_type().is_dir() => metadata,
110 Ok(_) => continue,
111 Err(error) if error.kind() == std::io::ErrorKind::NotFound => continue,
112 Err(error) => {
113 return Err(Error::new(format!(
114 "cannot inspect {}: {error}",
115 context_path.display()
116 )));
117 }
118 };
119 if !context_metadata.is_dir() {
120 continue;
121 }
122
123 candidates.push(ContextPackage {
124 name: package.name,
125 version: package.version,
126 description: package.description,
127 context_path,
128 selector: String::new(),
129 });
130 }
131
132 let mut name_counts = HashMap::new();
133 for package in &candidates {
134 *name_counts.entry(package.name.clone()).or_insert(0usize) += 1;
135 }
136 for package in &mut candidates {
137 package.selector = if name_counts[&package.name] == 1 {
138 package.name.clone()
139 } else {
140 format!("{}@{}", package.name, package.version)
141 };
142 }
143 candidates.sort_by(|left, right| {
144 left.name
145 .cmp(&right.name)
146 .then_with(|| left.version.cmp(&right.version))
147 });
148
149 Ok(Self {
150 context_path: root.join(".agents/context"),
151 root,
152 packages: candidates,
153 })
154 }
155
156 pub fn root(&self) -> &Path {
157 &self.root
158 }
159
160 pub fn context_path(&self) -> &Path {
161 &self.context_path
162 }
163
164 pub fn packages(&self) -> &[ContextPackage] {
165 &self.packages
166 }
167
168 pub fn find_package(&self, selector: &str) -> Result<Option<ContextPackage>> {
169 let matches: Vec<_> = self
170 .packages
171 .iter()
172 .filter(|package| package.selector == selector || package.name == selector)
173 .collect();
174
175 match matches.as_slice() {
176 [] => Ok(None),
177 [package] => Ok(Some((*package).clone())),
178 _ => {
179 let selectors = matches
180 .iter()
181 .map(|package| package.selector.as_str())
182 .collect::<Vec<_>>()
183 .join(", ");
184 Err(Error::new(format!(
185 "multiple versions of crate {selector:?} provide context; choose one of: {selectors}"
186 )))
187 }
188 }
189 }
190
191 pub fn list_context_files(&self, package: &ContextPackage) -> Result<Vec<ContextFile>> {
192 let skill_names: HashSet<_> = super::skill::list_package_skills(package)?
193 .into_iter()
194 .map(|skill| skill.source_name)
195 .collect();
196 let mut files = Vec::new();
197 collect_files(&package.context_path, &mut files)?;
198 files.sort();
199 Ok(files
200 .into_iter()
201 .filter_map(|file| {
202 file.strip_prefix(&package.context_path)
203 .ok()
204 .map(|path| ContextFile {
205 path: path.to_path_buf(),
206 })
207 })
208 .filter(|file| !is_skill_context_path(&file.path, &skill_names))
209 .collect())
210 }
211
212 pub fn show_context_file(&self, selector: &str, file: &str) -> Result<Option<String>> {
213 let Some(package) = self.find_package(selector)? else {
214 return Ok(None);
215 };
216 let Some(path) = find_context_file(&package.context_path, file)? else {
217 return Ok(None);
218 };
219
220 let skill_names: HashSet<_> = super::skill::list_package_skills(&package)?
221 .into_iter()
222 .map(|skill| skill.source_name)
223 .collect();
224 let context_root = package.context_path.canonicalize().map_err(|error| {
225 Error::new(format!(
226 "cannot resolve {}: {error}",
227 package.context_path.display()
228 ))
229 })?;
230 let relative_path = path
231 .strip_prefix(&context_root)
232 .map_err(|error| Error::new(format!("cannot make context path relative: {error}")))?;
233 if is_skill_context_path(relative_path, &skill_names) {
234 return Ok(None);
235 }
236
237 fs::read_to_string(&path)
238 .map(Some)
239 .map_err(|error| Error::new(format!("cannot read {}: {error}", path.display())))
240 }
241
242 pub fn install_package(&self, selector: &str) -> Result<bool> {
244 let Some(package) = self.find_package(selector)? else {
245 return Ok(false);
246 };
247 let skills = super::skill::list_package_skills(&package)?;
248 let skill_names: HashSet<_> = skills.into_iter().map(|skill| skill.source_name).collect();
249
250 fs::create_dir_all(&self.context_path).map_err(|error| {
251 Error::new(format!(
252 "cannot create {}: {error}",
253 self.context_path.display()
254 ))
255 })?;
256 let destination = self.context_path.join(&package.selector);
257 remove_existing(&destination)?;
258 let copied = copy_context_tree(&package.context_path, &destination, &skill_names, true)?;
259 if !copied {
260 remove_existing(&destination)?;
261 }
262 Ok(copied)
263 }
264
265 pub fn install_all(&self) -> Result<Vec<String>> {
267 let mut installed = Vec::new();
268 for package in &self.packages {
269 if self.install_package(&package.selector)? {
270 installed.push(package.selector.clone());
271 }
272 }
273 Ok(installed)
274 }
275}
276
277#[derive(Deserialize)]
278struct CargoMetadata {
279 workspace_members: Vec<String>,
280 packages: Vec<CargoPackage>,
281 resolve: Option<Resolve>,
282}
283
284#[derive(Deserialize)]
285struct Resolve {
286 nodes: Vec<ResolveNode>,
287}
288
289#[derive(Deserialize)]
290struct ResolveNode {
291 #[serde(rename = "id")]
292 package_id: String,
293}
294
295#[derive(Deserialize)]
296struct CargoPackage {
297 #[serde(rename = "id")]
298 package_id: String,
299 name: String,
300 version: String,
301 description: Option<String>,
302 manifest_path: PathBuf,
303}
304
305fn find_context_file(context_path: &Path, file: &str) -> Result<Option<PathBuf>> {
306 let requested = Path::new(file);
307 if requested.is_absolute()
308 || requested
309 .components()
310 .any(|component| !matches!(component, Component::Normal(_)))
311 {
312 return Err(Error::new(
313 "context file must be a relative path inside context/",
314 ));
315 }
316
317 let mut candidates = vec![context_path.join(requested)];
318 if requested.extension().is_none() {
319 candidates.push(context_path.join(requested).with_extension("md"));
320 }
321
322 for candidate in candidates {
323 let Ok(canonical_candidate) = candidate.canonicalize() else {
324 continue;
325 };
326 let canonical_root = context_path.canonicalize().map_err(|error| {
327 Error::new(format!(
328 "cannot resolve {}: {error}",
329 context_path.display()
330 ))
331 })?;
332 if !canonical_candidate.starts_with(&canonical_root) || !canonical_candidate.is_file() {
333 continue;
334 }
335 return Ok(Some(canonical_candidate));
336 }
337
338 Ok(None)
339}
340
341pub(crate) fn markdown_files(root: &Path) -> Result<Vec<PathBuf>> {
342 let mut files = Vec::new();
343 collect_markdown_files(root, &mut files)?;
344 files.sort();
345 Ok(files)
346}
347
348fn collect_markdown_files(directory: &Path, files: &mut Vec<PathBuf>) -> Result<()> {
349 let entries = fs::read_dir(directory)
350 .map_err(|error| Error::new(format!("cannot read {}: {error}", directory.display())))?;
351 for entry in entries {
352 let entry = entry?;
353 let file_type = entry.file_type()?;
354 let path = entry.path();
355 if file_type.is_dir() {
356 collect_markdown_files(&path, files)?;
357 } else if file_type.is_file()
358 && path
359 .extension()
360 .is_some_and(|extension| extension.eq_ignore_ascii_case("md"))
361 {
362 files.push(path);
363 }
364 }
365 Ok(())
366}
367
368fn collect_files(directory: &Path, files: &mut Vec<PathBuf>) -> Result<()> {
369 let entries = fs::read_dir(directory)
370 .map_err(|error| Error::new(format!("cannot read {}: {error}", directory.display())))?;
371 for entry in entries {
372 let entry = entry?;
373 let file_type = entry.file_type()?;
374 let path = entry.path();
375 if file_type.is_dir() {
376 collect_files(&path, files)?;
377 } else if file_type.is_file() {
378 files.push(path);
379 }
380 }
381 Ok(())
382}
383
384fn copy_context_tree(
385 source: &Path,
386 destination: &Path,
387 skill_names: &HashSet<String>,
388 root: bool,
389) -> Result<bool> {
390 let source_type = fs::symlink_metadata(source)
391 .map_err(|error| Error::new(format!("cannot inspect {}: {error}", source.display())))?
392 .file_type();
393 if !source_type.is_dir() {
394 return Err(Error::new(format!(
395 "context provider {} is not a regular directory",
396 source.display()
397 )));
398 }
399
400 fs::create_dir_all(destination)
401 .map_err(|error| Error::new(format!("cannot create {}: {error}", destination.display())))?;
402 let entries = fs::read_dir(source)
403 .map_err(|error| Error::new(format!("cannot read {}: {error}", source.display())))?;
404 let mut copied = false;
405 for entry in entries {
406 let entry = entry?;
407 let file_type = entry.file_type()?;
408 let source_path = entry.path();
409 let destination_path = destination.join(entry.file_name());
410 if file_type.is_dir() {
411 if root
412 && entry
413 .file_name()
414 .to_str()
415 .is_some_and(|name| skill_names.contains(name))
416 {
417 continue;
418 }
419
420 if copy_context_tree(&source_path, &destination_path, skill_names, false)? {
421 copied = true;
422 } else {
423 fs::remove_dir(&destination_path).map_err(|error| {
424 Error::new(format!(
425 "cannot remove empty context directory {}: {error}",
426 destination_path.display()
427 ))
428 })?;
429 }
430 } else if file_type.is_file() {
431 if root && is_skill_markdown(&source_path, skill_names) {
432 continue;
433 }
434 fs::copy(&source_path, &destination_path).map_err(|error| {
435 Error::new(format!(
436 "cannot copy {} to {}: {error}",
437 source_path.display(),
438 destination_path.display()
439 ))
440 })?;
441 copied = true;
442 }
443 }
444 Ok(copied)
445}
446
447fn is_skill_context_path(path: &Path, skill_names: &HashSet<String>) -> bool {
448 let mut components = path.components();
449 let Some(first) = components
450 .next()
451 .and_then(|component| component.as_os_str().to_str())
452 else {
453 return false;
454 };
455 if skill_names.contains(first) {
456 return true;
457 }
458
459 components.next().is_none() && is_skill_markdown(path, skill_names)
460}
461
462fn is_skill_markdown(path: &Path, skill_names: &HashSet<String>) -> bool {
463 path.extension()
464 .is_some_and(|extension| extension.eq_ignore_ascii_case("md"))
465 && path
466 .file_stem()
467 .and_then(|stem| stem.to_str())
468 .is_some_and(|stem| skill_names.contains(stem))
469}
470
471fn remove_existing(path: &Path) -> Result<()> {
472 let metadata = match fs::symlink_metadata(path) {
473 Ok(metadata) => metadata,
474 Err(error) if error.kind() == std::io::ErrorKind::NotFound => return Ok(()),
475 Err(error) => {
476 return Err(Error::new(format!(
477 "cannot inspect {}: {error}",
478 path.display()
479 )));
480 }
481 };
482 let result = if metadata.file_type().is_dir() {
483 fs::remove_dir_all(path)
484 } else {
485 fs::remove_file(path)
486 };
487 result.map_err(|error| Error::new(format!("cannot remove {}: {error}", path.display())))
488}