Skip to main content

bake_agent_context/agent/context/
installer.rs

1// Released under the MIT License.
2// Copyright, 2026, by Samuel Williams.
3
4use 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/// A resolved Cargo package that provides a top-level `context/` directory.
13#[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    /// The name accepted by the task's `--package` option.
24    pub fn selector(&self) -> &str {
25        &self.selector
26    }
27}
28
29/// A context file relative to its provider's `context/` directory.
30#[derive(Clone, Debug, Eq, PartialEq)]
31pub struct ContextFile {
32    pub path: PathBuf,
33}
34
35/// Discovers, lists, reads, and installs context from resolved Cargo dependencies.
36#[derive(Clone, Debug)]
37pub struct Installer {
38    root: PathBuf,
39    context_path: PathBuf,
40    packages: Vec<ContextPackage>,
41}
42
43impl Installer {
44    /// Resolve the project's Cargo packages and find dependencies with `context/` directories.
45    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    /// Install one package's context. Returns `false` when it does not provide context.
243    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    /// Install all resolved dependency packages that provide context.
266    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}