Skip to main content

ag_harness/
harness.rs

1use std::num::NonZeroUsize;
2use std::path::PathBuf;
3use std::sync::Arc;
4
5use serde_json::Value;
6use thiserror::Error;
7
8use crate::file_system::{FileSystem, LocalFileSystem};
9use crate::model::{Model, ModelError, ModelRequest, ModelResponse};
10use crate::policy::Policy;
11use crate::read::{ReadError, ReadTool};
12use crate::schema_contract::OutputSchema;
13use crate::tool::{Tool, ToolDefinition};
14
15const DEFAULT_MAX_TOOL_CALLS: usize = 8;
16
17/// Application-facing harness for one complete model turn.
18///
19/// A turn advertises policy-approved tools, executes validated native calls,
20/// returns tool results to the model, and finishes with locally validated
21/// structured output.
22pub struct Harness {
23    file_system: Arc<dyn FileSystem>,
24    max_tool_calls: usize,
25    model: Arc<dyn Model>,
26    policy: Policy,
27    repository_root: Option<PathBuf>,
28}
29
30impl Harness {
31    /// Creates a deny-by-default harness backed by the local filesystem.
32    pub fn new(model: impl Model + 'static) -> Self {
33        Self {
34            file_system: Arc::new(LocalFileSystem),
35            max_tool_calls: DEFAULT_MAX_TOOL_CALLS,
36            model: Arc::new(model),
37            policy: Policy::default(),
38            repository_root: None,
39        }
40    }
41
42    /// Roots repository-scoped tools at `repository_root`.
43    #[must_use]
44    pub fn repository(mut self, repository_root: impl Into<PathBuf>) -> Self {
45        self.repository_root = Some(repository_root.into());
46
47        self
48    }
49
50    /// Enables one built-in tool for model requests.
51    #[must_use]
52    pub fn allow(mut self, tool: Tool) -> Self {
53        self.policy.allow(tool);
54
55        self
56    }
57
58    /// Replaces the local filesystem implementation.
59    #[must_use]
60    pub fn file_system(mut self, file_system: impl FileSystem + 'static) -> Self {
61        self.file_system = Arc::new(file_system);
62
63        self
64    }
65
66    /// Overrides the maximum number of native calls allowed in one turn.
67    #[must_use]
68    pub fn max_tool_calls(mut self, max_tool_calls: NonZeroUsize) -> Self {
69        self.max_tool_calls = max_tool_calls.get();
70
71        self
72    }
73
74    /// Runs one prompt through tool execution to terminal structured output.
75    ///
76    /// # Errors
77    ///
78    /// Returns [`TurnError`] when the model fails, requests a denied tool,
79    /// exceeds the call limit, or the requested repository read fails.
80    pub async fn run(
81        &self,
82        prompt: impl Into<String>,
83        schema: OutputSchema,
84    ) -> Result<Value, TurnError> {
85        let mut request = ModelRequest::new(prompt, schema);
86        let read_tool = if self.policy.allows(Tool::Read) {
87            let repository_root = self
88                .repository_root
89                .as_ref()
90                .ok_or(TurnError::RepositoryRequired)?;
91            request = request.with_tool(ToolDefinition::read());
92
93            Some(ReadTool::new(
94                self.file_system.clone(),
95                repository_root.clone(),
96            ))
97        } else {
98            None
99        };
100        let mut completed_tool_calls = 0_usize;
101
102        loop {
103            match self.model.complete(request.clone()).await? {
104                ModelResponse::Output(output) => return Ok(output),
105                ModelResponse::ToolCall(call) => {
106                    let Some(read_tool) = read_tool.as_ref() else {
107                        return Err(TurnError::ToolDenied {
108                            name: call.name().to_string(),
109                        });
110                    };
111                    if completed_tool_calls >= self.max_tool_calls {
112                        return Err(TurnError::ToolCallLimit {
113                            limit: self.max_tool_calls,
114                        });
115                    }
116                    let output = read_tool.execute(call.arguments()).await?;
117                    let result = output.to_tool_result().map_err(ReadError::from)?;
118                    request.record_tool_result(call, result);
119                    completed_tool_calls += 1;
120                }
121            }
122        }
123    }
124}
125
126/// Failure returned by a complete harness turn.
127#[derive(Debug, Error)]
128pub enum TurnError {
129    /// Provider request, response decoding, or terminal validation failed.
130    #[error(transparent)]
131    Model(#[from] ModelError),
132    /// The model requested a tool unavailable under the configured policy.
133    #[error("tool `{name}` is denied by policy")]
134    ToolDenied {
135        /// Denied native function name.
136        name: String,
137    },
138    /// A repository read failed.
139    #[error(transparent)]
140    Read(#[from] ReadError),
141    /// Repository-scoped tools were enabled without a repository root.
142    #[error("repository root is required when the read tool is allowed")]
143    RepositoryRequired,
144    /// The model exceeded the bounded number of calls in one turn.
145    #[error("model exceeded the per-turn tool call limit of {limit}")]
146    ToolCallLimit {
147        /// Configured maximum calls.
148        limit: usize,
149    },
150}
151
152#[cfg(test)]
153mod tests {
154    use std::io::{self, Cursor};
155    use std::sync::atomic::{AtomicUsize, Ordering};
156
157    use mockall::Sequence;
158    use serde_json::json;
159
160    use super::*;
161    use crate::file_system::MockFileSystem;
162    use crate::model::{MockModel, ModelMessage};
163    use crate::tool::{ReadArguments, ToolCall};
164
165    fn object_schema() -> OutputSchema {
166        OutputSchema::new(json!({
167            "type": "object",
168            "properties": { "summary": { "type": "string" } },
169            "required": ["summary"],
170            "additionalProperties": false
171        }))
172        .expect("schema should be valid")
173    }
174
175    fn read_call(id: &str) -> ToolCall {
176        let arguments = serde_json::from_value::<ReadArguments>(json!({
177            "path": "Cargo.toml",
178            "limit": 1
179        }))
180        .expect("read arguments should be valid");
181
182        ToolCall::read(id.to_string(), arguments, None)
183    }
184
185    fn readable_file_system() -> MockFileSystem {
186        let mut file_system = MockFileSystem::new();
187        let mut sequence = Sequence::new();
188        file_system
189            .expect_canonicalize()
190            .times(1)
191            .in_sequence(&mut sequence)
192            .returning(|_| Ok(PathBuf::from("/repo")));
193        file_system
194            .expect_canonicalize()
195            .times(1)
196            .in_sequence(&mut sequence)
197            .returning(|_| Ok(PathBuf::from("/repo/Cargo.toml")));
198        file_system
199            .expect_open_beneath()
200            .times(1)
201            .returning(|_, _| {
202                Ok(Box::new(Cursor::new(
203                    b"[workspace]\nmember = true\n".to_vec(),
204                )))
205            });
206
207        file_system
208    }
209
210    #[tokio::test]
211    async fn completes_read_tool_round_trip() {
212        // Arrange
213        let mut model = MockModel::new();
214        let call_count = Arc::new(AtomicUsize::new(0));
215        model.expect_complete().times(2).returning(move |request| {
216            let call_index = call_count.fetch_add(1, Ordering::SeqCst);
217            if call_index == 0 {
218                assert_eq!(request.tools(), &[ToolDefinition::read()]);
219
220                return Ok(ModelResponse::ToolCall(read_call("call_read")));
221            }
222            assert_eq!(request.messages().len(), 3);
223            assert!(matches!(
224                &request.messages()[0],
225                ModelMessage::User(prompt) if prompt == "inspect the manifest"
226            ));
227            assert!(matches!(
228                &request.messages()[1],
229                ModelMessage::AssistantToolCall(call) if call.id() == "call_read"
230            ));
231            assert!(matches!(
232                &request.messages()[2],
233                ModelMessage::ToolResult {
234                    call_id,
235                    content,
236                    name,
237                }
238                    if call_id == "call_read"
239                        && name == "read"
240                        && serde_json::from_str::<Value>(content)
241                            .is_ok_and(|value| value["content"] == "[workspace]")
242            ));
243
244            Ok(ModelResponse::Output(json!({ "summary": "workspace" })))
245        });
246        let harness = Harness::new(model)
247            .file_system(readable_file_system())
248            .repository("repo")
249            .allow(Tool::Read);
250
251        // Act
252        let output = harness
253            .run("inspect the manifest", object_schema())
254            .await
255            .expect("tool round trip should succeed");
256
257        // Assert
258        assert_eq!(output, json!({ "summary": "workspace" }));
259    }
260
261    #[tokio::test]
262    async fn requires_repository_when_read_is_allowed() {
263        // Arrange
264        let mut model = MockModel::new();
265        model.expect_complete().times(0);
266        let harness = Harness::new(model).allow(Tool::Read);
267
268        // Act
269        let error = harness
270            .run("inspect", object_schema())
271            .await
272            .expect_err("read should require a repository root");
273
274        // Assert
275        assert!(matches!(error, TurnError::RepositoryRequired));
276    }
277
278    #[tokio::test]
279    async fn rejects_tool_call_when_policy_denies_read() {
280        // Arrange
281        let mut model = MockModel::new();
282        model.expect_complete().times(1).returning(|request| {
283            assert!(request.tools().is_empty());
284
285            Ok(ModelResponse::ToolCall(read_call("call_denied")))
286        });
287        let mut file_system = MockFileSystem::new();
288        file_system.expect_canonicalize().times(0);
289        file_system.expect_open_beneath().times(0);
290        let harness = Harness::new(model)
291            .file_system(file_system)
292            .repository("repo");
293
294        // Act
295        let error = harness
296            .run("inspect", object_schema())
297            .await
298            .expect_err("denied tool should fail");
299
300        // Assert
301        assert!(matches!(
302            error,
303            TurnError::ToolDenied { name } if name == "read"
304        ));
305    }
306
307    #[tokio::test]
308    async fn enforces_tool_call_limit() {
309        // Arrange
310        let mut model = MockModel::new();
311        model
312            .expect_complete()
313            .times(2)
314            .returning(|_| Ok(ModelResponse::ToolCall(read_call("call_read"))));
315        let harness = Harness::new(model)
316            .file_system(readable_file_system())
317            .repository("repo")
318            .allow(Tool::Read)
319            .max_tool_calls(NonZeroUsize::new(1).expect("limit should be non-zero"));
320
321        // Act
322        let error = harness
323            .run("inspect", object_schema())
324            .await
325            .expect_err("second tool call should exceed the limit");
326
327        // Assert
328        assert!(matches!(error, TurnError::ToolCallLimit { limit: 1 }));
329    }
330
331    #[tokio::test]
332    async fn returns_typed_read_failure() {
333        // Arrange
334        let mut model = MockModel::new();
335        model
336            .expect_complete()
337            .times(1)
338            .returning(|_| Ok(ModelResponse::ToolCall(read_call("call_read"))));
339        let mut file_system = MockFileSystem::new();
340        file_system
341            .expect_canonicalize()
342            .times(1)
343            .returning(|_| Err(io::Error::new(io::ErrorKind::NotFound, "missing root")));
344        let harness = Harness::new(model)
345            .file_system(file_system)
346            .repository("repo")
347            .allow(Tool::Read);
348
349        // Act
350        let error = harness
351            .run("inspect", object_schema())
352            .await
353            .expect_err("filesystem failure should end the turn");
354
355        // Assert
356        assert!(matches!(
357            error,
358            TurnError::Read(ReadError::RepositoryRoot { .. })
359        ));
360    }
361
362    #[tokio::test]
363    async fn returns_typed_model_failure() {
364        // Arrange
365        let mut model = MockModel::new();
366        model
367            .expect_complete()
368            .times(1)
369            .returning(|_| Err(ModelError::request(io::Error::other("offline"))));
370        let mut file_system = MockFileSystem::new();
371        file_system.expect_canonicalize().times(0);
372        file_system.expect_open_beneath().times(0);
373        let harness = Harness::new(model)
374            .file_system(file_system)
375            .repository("repo");
376
377        // Act
378        let error = harness
379            .run("inspect", object_schema())
380            .await
381            .expect_err("model failure should end the turn");
382
383        // Assert
384        assert!(matches!(error, TurnError::Model(ModelError::Request(_))));
385    }
386}