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 mut files = Vec::new();
193        collect_files(&package.context_path, &mut files)?;
194        files.sort();
195        Ok(files
196            .into_iter()
197            .filter_map(|file| {
198                file.strip_prefix(&package.context_path)
199                    .ok()
200                    .map(|path| ContextFile {
201                        path: path.to_path_buf(),
202                    })
203            })
204            .collect())
205    }
206
207    pub fn show_context_file(&self, selector: &str, file: &str) -> Result<Option<String>> {
208        let Some(package) = self.find_package(selector)? else {
209            return Ok(None);
210        };
211        let Some(path) = find_context_file(&package.context_path, file)? else {
212            return Ok(None);
213        };
214        fs::read_to_string(&path)
215            .map(Some)
216            .map_err(|error| Error::new(format!("cannot read {}: {error}", path.display())))
217    }
218
219    /// Install one package's context. Returns `false` when it does not provide context.
220    pub fn install_package(&self, selector: &str) -> Result<bool> {
221        let Some(package) = self.find_package(selector)? else {
222            return Ok(false);
223        };
224
225        fs::create_dir_all(&self.context_path).map_err(|error| {
226            Error::new(format!(
227                "cannot create {}: {error}",
228                self.context_path.display()
229            ))
230        })?;
231        let destination = self.context_path.join(&package.selector);
232        remove_existing(&destination)?;
233        copy_context_tree(&package.context_path, &destination)?;
234        Ok(true)
235    }
236
237    /// Install all resolved dependency packages that provide context.
238    pub fn install_all(&self) -> Result<Vec<String>> {
239        let mut installed = Vec::new();
240        for package in &self.packages {
241            if self.install_package(&package.selector)? {
242                installed.push(package.selector.clone());
243            }
244        }
245        Ok(installed)
246    }
247}
248
249#[derive(Deserialize)]
250struct CargoMetadata {
251    workspace_members: Vec<String>,
252    packages: Vec<CargoPackage>,
253    resolve: Option<Resolve>,
254}
255
256#[derive(Deserialize)]
257struct Resolve {
258    nodes: Vec<ResolveNode>,
259}
260
261#[derive(Deserialize)]
262struct ResolveNode {
263    #[serde(rename = "id")]
264    package_id: String,
265}
266
267#[derive(Deserialize)]
268struct CargoPackage {
269    #[serde(rename = "id")]
270    package_id: String,
271    name: String,
272    version: String,
273    description: Option<String>,
274    manifest_path: PathBuf,
275}
276
277fn find_context_file(context_path: &Path, file: &str) -> Result<Option<PathBuf>> {
278    let requested = Path::new(file);
279    if requested.is_absolute()
280        || requested
281            .components()
282            .any(|component| !matches!(component, Component::Normal(_)))
283    {
284        return Err(Error::new(
285            "context file must be a relative path inside context/",
286        ));
287    }
288
289    let mut candidates = vec![context_path.join(requested)];
290    if requested.extension().is_none() {
291        candidates.push(context_path.join(requested).with_extension("md"));
292    }
293
294    for candidate in candidates {
295        let Ok(canonical_candidate) = candidate.canonicalize() else {
296            continue;
297        };
298        let canonical_root = context_path.canonicalize().map_err(|error| {
299            Error::new(format!(
300                "cannot resolve {}: {error}",
301                context_path.display()
302            ))
303        })?;
304        if !canonical_candidate.starts_with(&canonical_root) || !canonical_candidate.is_file() {
305            continue;
306        }
307        return Ok(Some(canonical_candidate));
308    }
309
310    Ok(None)
311}
312
313pub(crate) fn markdown_files(root: &Path) -> Result<Vec<PathBuf>> {
314    let mut files = Vec::new();
315    collect_markdown_files(root, &mut files)?;
316    files.sort();
317    Ok(files)
318}
319
320fn collect_markdown_files(directory: &Path, files: &mut Vec<PathBuf>) -> Result<()> {
321    let entries = fs::read_dir(directory)
322        .map_err(|error| Error::new(format!("cannot read {}: {error}", directory.display())))?;
323    for entry in entries {
324        let entry = entry?;
325        let file_type = entry.file_type()?;
326        let path = entry.path();
327        if file_type.is_dir() {
328            collect_markdown_files(&path, files)?;
329        } else if file_type.is_file()
330            && path
331                .extension()
332                .is_some_and(|extension| extension.eq_ignore_ascii_case("md"))
333        {
334            files.push(path);
335        }
336    }
337    Ok(())
338}
339
340fn collect_files(directory: &Path, files: &mut Vec<PathBuf>) -> Result<()> {
341    let entries = fs::read_dir(directory)
342        .map_err(|error| Error::new(format!("cannot read {}: {error}", directory.display())))?;
343    for entry in entries {
344        let entry = entry?;
345        let file_type = entry.file_type()?;
346        let path = entry.path();
347        if file_type.is_dir() {
348            collect_files(&path, files)?;
349        } else if file_type.is_file() {
350            files.push(path);
351        }
352    }
353    Ok(())
354}
355
356fn copy_context_tree(source: &Path, destination: &Path) -> Result<()> {
357    let source_type = fs::symlink_metadata(source)
358        .map_err(|error| Error::new(format!("cannot inspect {}: {error}", source.display())))?
359        .file_type();
360    if !source_type.is_dir() {
361        return Err(Error::new(format!(
362            "context provider {} is not a regular directory",
363            source.display()
364        )));
365    }
366
367    fs::create_dir_all(destination)
368        .map_err(|error| Error::new(format!("cannot create {}: {error}", destination.display())))?;
369    let entries = fs::read_dir(source)
370        .map_err(|error| Error::new(format!("cannot read {}: {error}", source.display())))?;
371    for entry in entries {
372        let entry = entry?;
373        let file_type = entry.file_type()?;
374        let source_path = entry.path();
375        let destination_path = destination.join(entry.file_name());
376        if file_type.is_dir() {
377            copy_context_tree(&source_path, &destination_path)?;
378        } else if file_type.is_file() {
379            fs::copy(&source_path, &destination_path).map_err(|error| {
380                Error::new(format!(
381                    "cannot copy {} to {}: {error}",
382                    source_path.display(),
383                    destination_path.display()
384                ))
385            })?;
386        }
387    }
388    Ok(())
389}
390
391fn remove_existing(path: &Path) -> Result<()> {
392    let metadata = match fs::symlink_metadata(path) {
393        Ok(metadata) => metadata,
394        Err(error) if error.kind() == std::io::ErrorKind::NotFound => return Ok(()),
395        Err(error) => {
396            return Err(Error::new(format!(
397                "cannot inspect {}: {error}",
398                path.display()
399            )));
400        }
401    };
402    let result = if metadata.file_type().is_dir() {
403        fs::remove_dir_all(path)
404    } else {
405        fs::remove_file(path)
406    };
407    result.map_err(|error| Error::new(format!("cannot remove {}: {error}", path.display())))
408}