Skip to main content

aether_evals/evals/
workspace.rs

1use super::diff::{DiffStats, GitDiff};
2use crate::WorkspaceError;
3use crate::git_repo::{CloneMode, GitRepo};
4use schemars::JsonSchema;
5use serde::Serialize;
6use std::{
7    fs::{create_dir_all, write},
8    path::{Path, PathBuf},
9};
10use tempfile::TempDir;
11
12pub struct Workspace {
13    path: PathBuf,
14    root_path: PathBuf,
15    relative_cwd: Option<PathBuf>,
16    source: WorkspaceSource,
17    temp_dir: TempDir,
18}
19
20#[derive(Debug, Clone)]
21pub enum WorkspaceSource {
22    Local,
23    GitRepo { url: String, start_commit: String, gold_commit: String },
24    Bundle { start_commit: String, gold_commit: String },
25}
26
27#[derive(Debug, Clone, serde::Deserialize, JsonSchema)]
28#[serde(rename_all = "camelCase", deny_unknown_fields)]
29pub struct GitRepoSpec {
30    pub url: String,
31    pub start_commit: String,
32    pub gold_commit: String,
33    #[serde(default)]
34    pub subdir: Option<PathBuf>,
35}
36
37#[derive(Debug, Clone, serde::Deserialize, JsonSchema)]
38#[serde(rename_all = "camelCase", deny_unknown_fields)]
39pub struct GitBundleSpec {
40    pub bundle_path: PathBuf,
41    pub start_commit: String,
42    pub gold_commit: String,
43    #[serde(default)]
44    pub subdir: Option<PathBuf>,
45}
46
47#[derive(Debug, Clone, Serialize, JsonSchema)]
48#[serde(rename_all = "camelCase", deny_unknown_fields)]
49pub struct RetainedWorkspaceInfo {
50    pub root_path: PathBuf,
51    pub path: PathBuf,
52}
53
54impl Workspace {
55    pub fn empty() -> Result<Self, WorkspaceError> {
56        let temp_dir = new_temp_dir()?;
57        let path = temp_dir.path().to_path_buf();
58        Ok(Self { path: path.clone(), root_path: path, relative_cwd: None, source: WorkspaceSource::Local, temp_dir })
59    }
60
61    pub fn from_dir(src_path: impl Into<PathBuf>) -> Result<Self, WorkspaceError> {
62        let src_path = src_path.into();
63        let temp_dir = new_temp_dir()?;
64        let path = temp_dir.path().to_path_buf();
65
66        copy_dir_contents(&src_path, &path).map_err(|source| WorkspaceError::CopyFixture {
67            from: src_path.clone(),
68            to: path.clone(),
69            source,
70        })?;
71
72        Ok(Self { path: path.clone(), root_path: path, relative_cwd: None, source: WorkspaceSource::Local, temp_dir })
73    }
74
75    pub fn from_files<T: AsRef<Path>, U: AsRef<str>>(
76        files: impl IntoIterator<Item = (T, U)>,
77    ) -> Result<Self, WorkspaceError> {
78        let workspace = Self::empty()?;
79        for (relative_path, contents) in files {
80            let path = workspace.path().join(relative_path.as_ref());
81            if let Some(parent) = path.parent() {
82                create_dir_all(parent)
83                    .map_err(|source| WorkspaceError::WriteFile { path: parent.to_path_buf(), source })?;
84            }
85            write(&path, contents.as_ref()).map_err(|source| WorkspaceError::WriteFile { path, source })?;
86        }
87        Ok(workspace)
88    }
89
90    pub fn from_git_repo(spec: GitRepoSpec) -> Result<Self, WorkspaceError> {
91        let GitRepoSpec { url, start_commit, gold_commit, subdir } = spec;
92        let temp_dir = new_temp_dir()?;
93
94        tracing::debug!("Cloning git repo {} at commit {}", url, start_commit);
95        let repo = GitRepo::clone(&url, temp_dir.path(), CloneMode::Blobless)?;
96
97        let source = WorkspaceSource::GitRepo { url, start_commit: start_commit.clone(), gold_commit };
98        Self::from_cloned(temp_dir, &repo, &start_commit, subdir, source)
99    }
100
101    /// Instantiate a workspace from a local git bundle file
102    pub fn from_git_bundle(spec: GitBundleSpec) -> Result<Self, WorkspaceError> {
103        let GitBundleSpec { bundle_path, start_commit, gold_commit, subdir } = spec;
104        if !bundle_path.exists() {
105            return Err(WorkspaceError::MissingBundle { path: bundle_path });
106        }
107
108        tracing::debug!("Cloning git bundle {} at commit {}", bundle_path.display(), start_commit);
109        let temp_dir = new_temp_dir()?;
110        let repo = GitRepo::clone(&bundle_path, temp_dir.path(), CloneMode::Full)?;
111        let source = WorkspaceSource::Bundle { start_commit: start_commit.clone(), gold_commit };
112        Self::from_cloned(temp_dir, &repo, &start_commit, subdir, source)
113    }
114
115    fn from_cloned(
116        temp_dir: TempDir,
117        repo: &GitRepo,
118        start_commit: &str,
119        subdir: Option<PathBuf>,
120        source: WorkspaceSource,
121    ) -> Result<Self, WorkspaceError> {
122        repo.checkout(start_commit)?;
123        let root_path = temp_dir.path().to_path_buf();
124        let (path, relative_cwd) = resolve_subdir(&root_path, subdir)?;
125        Ok(Self { path, root_path, relative_cwd, source, temp_dir })
126    }
127
128    pub fn path(&self) -> &Path {
129        &self.path
130    }
131
132    pub fn join(&self, relative_path: impl AsRef<Path>) -> PathBuf {
133        self.path.join(relative_path)
134    }
135
136    pub fn root_path(&self) -> &Path {
137        &self.root_path
138    }
139
140    pub fn relative_cwd(&self) -> Option<&Path> {
141        self.relative_cwd.as_deref()
142    }
143
144    pub fn source(&self) -> &WorkspaceSource {
145        &self.source
146    }
147
148    /// Prevents the workspace from getting automatically removed and returns its retained root and effective cwd. The caller is responsible for cleanup.
149    pub fn persist(self) -> RetainedWorkspaceInfo {
150        let root_path = self.temp_dir.keep();
151        let path = self.relative_cwd.map_or_else(|| root_path.clone(), |relative_cwd| root_path.join(relative_cwd));
152        RetainedWorkspaceInfo { root_path, path }
153    }
154
155    pub fn capture_git_diffs(&self) -> (Option<GitDiff>, Option<GitDiff>) {
156        let Some((start_commit, gold_commit)) = self.diff_commits() else {
157            return (None, None);
158        };
159
160        let repo = GitRepo::from_path(self.path());
161        let agent_diff =
162            repo.diff_range(start_commit, None).ok().map(|diff| GitDiff { stats: DiffStats::from_diff(&diff), diff });
163        let reference_diff =
164            repo.diff(start_commit, gold_commit).ok().map(|diff| GitDiff { stats: DiffStats::from_diff(&diff), diff });
165
166        (agent_diff, reference_diff)
167    }
168
169    fn diff_commits(&self) -> Option<(&str, &str)> {
170        match &self.source {
171            WorkspaceSource::Local => None,
172            WorkspaceSource::GitRepo { start_commit, gold_commit, .. }
173            | WorkspaceSource::Bundle { start_commit, gold_commit } => Some((start_commit, gold_commit)),
174        }
175    }
176}
177
178const EVAL_START_REF: &str = "eval-start";
179const EVAL_GOLD_REF: &str = "eval-gold";
180
181/// Create a self-contained git bundle at `out` containing `spec`'s start and gold commits.
182pub fn create_git_bundle(spec: &GitRepoSpec, out: &Path) -> Result<(), WorkspaceError> {
183    let temp_dir = new_temp_dir()?;
184    let repo = GitRepo::init(temp_dir.path())?;
185    repo.fetch(&spec.url, &[&spec.start_commit, &spec.gold_commit])?;
186    repo.update_ref(&format!("refs/heads/{EVAL_START_REF}"), &spec.start_commit)?;
187    repo.update_ref(&format!("refs/heads/{EVAL_GOLD_REF}"), &spec.gold_commit)?;
188    repo.bundle(&[EVAL_START_REF, EVAL_GOLD_REF], out)?;
189    Ok(())
190}
191
192fn new_temp_dir() -> Result<tempfile::TempDir, WorkspaceError> {
193    tempfile::tempdir().map_err(WorkspaceError::CreateTempDir)
194}
195
196fn resolve_subdir(root_path: &Path, subdir: Option<PathBuf>) -> Result<(PathBuf, Option<PathBuf>), WorkspaceError> {
197    match subdir {
198        None => Ok((root_path.to_path_buf(), None)),
199        Some(relative_cwd) => {
200            let working_path = root_path.join(&relative_cwd);
201            if !working_path.exists() {
202                return Err(WorkspaceError::MissingSubdir { path: working_path });
203            }
204            Ok((working_path, Some(relative_cwd)))
205        }
206    }
207}
208
209fn copy_dir_contents(src: &Path, dst: &Path) -> std::io::Result<()> {
210    for entry in std::fs::read_dir(src)? {
211        let entry = entry?;
212        let source_path = entry.path();
213        let dest_path = dst.join(entry.file_name());
214        let file_type = entry.file_type()?;
215
216        if file_type.is_dir() {
217            std::fs::create_dir_all(&dest_path)?;
218            copy_dir_contents(&source_path, &dest_path)?;
219        } else if file_type.is_file() {
220            std::fs::copy(&source_path, &dest_path)?;
221        }
222    }
223    Ok(())
224}
225
226#[cfg(test)]
227mod tests {
228    use super::*;
229    use std::fs::{read_to_string, remove_dir_all};
230    use std::process::Command;
231    use tempfile::TempDir;
232
233    #[test]
234    fn persist_reports_root_and_path_for_local_workspace() {
235        let workspace = Workspace::from_files([("notes.txt", "hi\n")]).unwrap();
236
237        let retained = workspace.persist();
238
239        assert_eq!(retained.root_path, retained.path);
240        assert_eq!(read_to_string(retained.path.join("notes.txt")).unwrap(), "hi\n");
241        remove_dir_all(retained.root_path).unwrap();
242    }
243
244    #[test]
245    fn from_git_bundle_round_trips_checkout_subdir_and_diffs() {
246        let (repo, start, gold) = init_repo();
247        let bundle_dir = tempfile::tempdir().unwrap();
248        let bundle_path = bundle_dir.path().join("repo.bundle");
249        create_git_bundle(
250            &GitRepoSpec {
251                url: format!("file://{}", repo.path().display()),
252                start_commit: start.clone(),
253                gold_commit: gold.clone(),
254                subdir: None,
255            },
256            &bundle_path,
257        )
258        .unwrap();
259
260        let workspace = Workspace::from_git_bundle(GitBundleSpec {
261            bundle_path,
262            start_commit: start,
263            gold_commit: gold,
264            subdir: Some(PathBuf::from("pkg")),
265        })
266        .unwrap();
267
268        let (agent_diff, reference_diff) = workspace.capture_git_diffs();
269
270        assert_eq!(read_to_string(workspace.root_path().join("root.txt")).unwrap(), "root v1\n");
271        assert_eq!(workspace.relative_cwd(), Some(Path::new("pkg")));
272        assert_eq!(workspace.path(), workspace.root_path().join("pkg"));
273        assert!(reference_diff.unwrap().diff.contains("root v2"), "reference diff should span start..gold");
274        assert!(agent_diff.unwrap().diff.is_empty(), "no agent edits yet");
275
276        write(workspace.root_path().join("root.txt"), "root edited\n").unwrap();
277        let (agent_diff, _) = workspace.capture_git_diffs();
278
279        assert!(agent_diff.unwrap().diff.contains("root edited"), "agent diff should capture the edit");
280    }
281
282    #[test]
283    fn capture_git_diffs_includes_committed_agent_changes() {
284        let (repo, start, gold) = init_repo();
285        let workspace = Workspace::from_git_repo(GitRepoSpec {
286            url: format!("file://{}", repo.path().display()),
287            start_commit: start,
288            gold_commit: gold,
289            subdir: None,
290        })
291        .unwrap();
292
293        write(workspace.join("root.txt"), "agent committed\n").unwrap();
294        git(workspace.root_path(), &["add", "."]);
295        git(
296            workspace.root_path(),
297            &["-c", "user.email=agent@example.com", "-c", "user.name=Agent", "commit", "-m", "agent change"],
298        );
299
300        let (agent_diff, _) = workspace.capture_git_diffs();
301        let agent_diff = agent_diff.unwrap();
302
303        assert!(agent_diff.diff.contains("+agent committed"));
304        assert_eq!(agent_diff.stats.files_changed, 1);
305        assert_eq!(agent_diff.stats.lines_added, 1);
306        assert_eq!(agent_diff.stats.lines_removed, 1);
307    }
308
309    #[test]
310    fn capture_git_diffs_includes_staged_unstaged_added_and_deleted_changes() {
311        let (repo, start, gold) = init_repo();
312        let workspace = Workspace::from_git_repo(GitRepoSpec {
313            url: format!("file://{}", repo.path().display()),
314            start_commit: start,
315            gold_commit: gold,
316            subdir: None,
317        })
318        .unwrap();
319
320        write(workspace.join("root.txt"), "staged change\n").unwrap();
321        write(workspace.join("added.txt"), "added\n").unwrap();
322        git(workspace.root_path(), &["add", "root.txt", "added.txt"]);
323        write(workspace.join("pkg/inner.txt"), "unstaged change\n").unwrap();
324        std::fs::remove_file(workspace.join("deleted.txt")).unwrap();
325
326        let (agent_diff, _) = workspace.capture_git_diffs();
327        let agent_diff = agent_diff.unwrap();
328
329        assert!(agent_diff.diff.contains("root.txt"));
330        assert!(agent_diff.diff.contains("+staged change"));
331        assert!(agent_diff.diff.contains("pkg/inner.txt"));
332        assert!(agent_diff.diff.contains("+unstaged change"));
333        assert!(agent_diff.diff.contains("added.txt"));
334        assert!(agent_diff.diff.contains("+added"));
335        assert!(agent_diff.diff.contains("deleted.txt"));
336        assert!(agent_diff.diff.contains("-before deleted"));
337        assert_eq!(agent_diff.stats.files_changed, 4);
338        assert_eq!(agent_diff.stats.lines_added, 3);
339        assert_eq!(agent_diff.stats.lines_removed, 3);
340    }
341
342    #[test]
343    fn from_git_bundle_missing_file_errors() {
344        let result = Workspace::from_git_bundle(GitBundleSpec {
345            bundle_path: PathBuf::from("/nonexistent/repo.bundle"),
346            start_commit: "abc".into(),
347            gold_commit: "def".into(),
348            subdir: None,
349        });
350
351        assert!(matches!(result, Err(WorkspaceError::MissingBundle { .. })));
352    }
353
354    fn init_repo() -> (TempDir, String, String) {
355        let dir = tempfile::tempdir().unwrap();
356        let path = dir.path();
357        git(path, &["init", "--initial-branch", "main"]);
358        git(path, &["config", "user.email", "eval@example.com"]);
359        git(path, &["config", "user.name", "Eval"]);
360
361        write(path.join("root.txt"), "root v1\n").unwrap();
362        write(path.join("deleted.txt"), "before deleted\n").unwrap();
363        create_dir_all(path.join("pkg")).unwrap();
364        write(path.join("pkg").join("inner.txt"), "inner v1\n").unwrap();
365        git(path, &["add", "."]);
366        git(path, &["commit", "-m", "start"]);
367        let start = git(path, &["rev-parse", "HEAD"]);
368
369        write(path.join("root.txt"), "root v2\n").unwrap();
370        git(path, &["add", "."]);
371        git(path, &["commit", "-m", "gold"]);
372        let gold = git(path, &["rev-parse", "HEAD"]);
373
374        (dir, start, gold)
375    }
376
377    fn git(repo: &Path, args: &[&str]) -> String {
378        let output = Command::new("git").arg("-C").arg(repo).args(args).output().unwrap();
379        assert!(output.status.success(), "git {args:?} failed: {}", String::from_utf8_lossy(&output.stderr));
380        String::from_utf8(output.stdout).unwrap().trim().to_string()
381    }
382}