Skip to repository content561 lines · 19.4 KB · rust
tenant.openagents/omega
No repository description is available.
OpenAgents Git authority 2026-07-28T06:44:09.169Z Public web read
NIP-34 coordinate
30617:7649603503856e5148d571eac2766b288a8ff1e9e35d380337a1d2b0015b4f92:omegaMaintainersHidden in public view
References2 branches · 1 tag
Read-only clone
git clone https://openagents.com/git/tenant.openagents/omega.gitBrowse files
write_file.rs
1use crate::{
2 AgentTool, ContextServerRegistry, ListDirectoryTool, ListDirectoryToolInput, Template,
3 Templates, Thread, ToolCallEventStream, ToolInput, WriteFileTool, WriteFileToolInput,
4};
5use Role::*;
6use anyhow::{Context as _, Result};
7use client::{Client, RefreshLlmTokenListener, UserStore};
8use fs::FakeFs;
9use futures::{FutureExt as _, StreamExt};
10use gpui::{AppContext as _, AsyncApp, Entity, TestAppContext, UpdateGlobal as _};
11use http_client::StatusCode;
12use language::language_settings::FormatOnSave;
13use language_model::{
14 LanguageModel, LanguageModelCompletionError, LanguageModelCompletionEvent,
15 LanguageModelRegistry, LanguageModelRequest, LanguageModelRequestMessage,
16 LanguageModelToolResult, LanguageModelToolResultContent, LanguageModelToolUse,
17 LanguageModelToolUseId, MessageContent, Role, SelectedModel,
18};
19use project::Project;
20use prompt_store::{ProjectContext, WorktreeContext};
21use rand::prelude::*;
22use reqwest_client::ReqwestClient;
23use serde::Serialize;
24use settings::SettingsStore;
25use std::{
26 fmt::{self, Display},
27 path::{Path, PathBuf},
28 str::FromStr,
29 sync::Arc,
30 time::Duration,
31};
32use util::path;
33
34#[derive(Clone)]
35struct EvalInput {
36 conversation: Vec<LanguageModelRequestMessage>,
37 input_file_path: PathBuf,
38 input_content: Option<String>,
39 expected_output_content: String,
40}
41
42impl EvalInput {
43 fn new(
44 conversation: Vec<LanguageModelRequestMessage>,
45 input_file_path: impl Into<PathBuf>,
46 input_content: Option<String>,
47 expected_output_content: String,
48 ) -> Self {
49 Self {
50 conversation,
51 input_file_path: input_file_path.into(),
52 input_content,
53 expected_output_content,
54 }
55 }
56}
57
58struct WriteEvalOutput {
59 tool_input: WriteFileToolInput,
60 text_after: String,
61}
62
63impl Display for WriteEvalOutput {
64 fn fmt(&self, f: &mut fmt::Formatter) -> fmt::Result {
65 writeln!(f, "Tool Input:\n{:#?}", self.tool_input)?;
66 writeln!(f, "Text After:\n{}", self.text_after)?;
67 Ok(())
68 }
69}
70
71struct WriteToolTest {
72 fs: Arc<FakeFs>,
73 project: Entity<Project>,
74 model: Arc<dyn LanguageModel>,
75 model_thinking_effort: Option<String>,
76}
77
78impl WriteToolTest {
79 async fn new(cx: &mut TestAppContext) -> Self {
80 cx.executor().allow_parking();
81
82 let fs = FakeFs::new(cx.executor());
83 cx.update(|cx| {
84 let settings_store = SettingsStore::test(cx);
85 cx.set_global(settings_store);
86 SettingsStore::update_global(cx, |store: &mut SettingsStore, cx| {
87 store.update_user_settings(cx, |settings| {
88 settings
89 .project
90 .all_languages
91 .defaults
92 .ensure_final_newline_on_save = Some(false);
93 settings.project.all_languages.defaults.format_on_save =
94 Some(FormatOnSave::Off);
95 });
96 });
97
98 gpui_tokio::init(cx);
99 let http_client = Arc::new(ReqwestClient::user_agent("agent tests").unwrap());
100 cx.set_http_client(http_client);
101 let client = Client::production(cx);
102 let user_store = cx.new(|cx| UserStore::new(client.clone(), cx));
103 language_model::init(cx);
104 RefreshLlmTokenListener::register(client.clone(), user_store.clone(), cx);
105 language_models::init(user_store, client, cx);
106 });
107
108 fs.insert_tree("/root", serde_json::json!({})).await;
109 let project = Project::test(fs.clone(), [path!("/root").as_ref()], cx).await;
110 let agent_model = SelectedModel::from_str(
111 &std::env::var("ZED_AGENT_MODEL")
112 .unwrap_or("anthropic/claude-sonnet-4-6-latest".into()),
113 )
114 .unwrap();
115
116 let authenticate_provider_tasks = cx.update(|cx| {
117 LanguageModelRegistry::global(cx).update(cx, |registry, cx| {
118 registry
119 .providers()
120 .iter()
121 .map(|p| p.authenticate(cx))
122 .collect::<Vec<_>>()
123 })
124 });
125 let model = cx
126 .update(|cx| {
127 cx.spawn(async move |cx| {
128 futures::future::join_all(authenticate_provider_tasks).await;
129 Self::load_model(&agent_model, cx).await.unwrap()
130 })
131 })
132 .await;
133
134 let model_thinking_effort = model
135 .default_effort_level()
136 .map(|effort_level| effort_level.value.to_string());
137
138 Self {
139 fs,
140 project,
141 model,
142 model_thinking_effort,
143 }
144 }
145
146 async fn load_model(
147 selected_model: &SelectedModel,
148 cx: &mut AsyncApp,
149 ) -> Result<Arc<dyn LanguageModel>> {
150 cx.update(|cx| {
151 let registry = LanguageModelRegistry::read_global(cx);
152 let provider = registry
153 .provider(&selected_model.provider)
154 .expect("Provider not found");
155 provider.authenticate(cx)
156 })
157 .await?;
158 Ok(cx.update(|cx| {
159 let models = LanguageModelRegistry::read_global(cx);
160 models
161 .available_models(cx)
162 .find(|model| {
163 model.provider_id() == selected_model.provider
164 && model.id() == selected_model.model
165 })
166 .unwrap_or_else(|| panic!("Model {} not found", selected_model.model.0))
167 }))
168 }
169
170 async fn eval(&self, mut eval: EvalInput, cx: &mut TestAppContext) -> Result<WriteEvalOutput> {
171 eval.conversation
172 .last_mut()
173 .context("Conversation must not be empty")?
174 .cache = true;
175
176 if let Some(input_content) = eval.input_content.as_deref() {
177 let abs_path = Path::new("/root").join(
178 eval.input_file_path
179 .strip_prefix("root")
180 .unwrap_or(&eval.input_file_path),
181 );
182 self.fs.insert_file(&abs_path, input_content.into()).await;
183 cx.run_until_parked();
184 }
185
186 let tools = crate::built_in_tools().collect::<Vec<_>>();
187
188 let system_prompt = {
189 let worktrees = vec![WorktreeContext {
190 root_name: "root".to_string(),
191 abs_path: Path::new("/path/to/root").into(),
192 rules_file: None,
193 }];
194 let project_context = ProjectContext::new(worktrees);
195 let tool_names = tools
196 .iter()
197 .map(|tool| tool.name.clone().into())
198 .collect::<Vec<_>>();
199 let template = crate::SystemPromptTemplate {
200 project: &project_context,
201 available_tools: tool_names,
202 model_name: None,
203 date: chrono::Local::now().format("%Y-%m-%d").to_string(),
204 user_agents_md: None,
205 sandboxing: false,
206 is_linux: cfg!(target_os = "linux"),
207 is_windows: cfg!(target_os = "windows"),
208 };
209 let templates = Templates::new();
210 template.render(&templates)?
211 };
212
213 let messages = [LanguageModelRequestMessage {
214 role: Role::System,
215 content: vec![MessageContent::Text(system_prompt)],
216 cache: true,
217 reasoning_details: None,
218 }]
219 .into_iter()
220 .chain(eval.conversation)
221 .collect::<Vec<_>>();
222
223 let request = LanguageModelRequest {
224 messages,
225 tools,
226 thinking_allowed: true,
227 thinking_effort: self.model_thinking_effort.clone(),
228 ..Default::default()
229 };
230
231 let tool_input =
232 retry_on_rate_limit(async || self.extract_tool_use(request.clone(), cx).await).await?;
233
234 let language_registry = self
235 .project
236 .read_with(cx, |project, _cx| project.languages().clone());
237
238 let context_server_registry = cx
239 .new(|cx| ContextServerRegistry::new(self.project.read(cx).context_server_store(), cx));
240 let thread = cx.new(|cx| {
241 Thread::new(
242 self.project.clone(),
243 cx.new(|_cx| ProjectContext::default()),
244 context_server_registry,
245 Templates::new(),
246 Some(self.model.clone()),
247 cx,
248 )
249 });
250 let action_log = thread.read_with(cx, |thread, _| thread.action_log().clone());
251
252 let tool = Arc::new(WriteFileTool::new(
253 self.project.clone(),
254 thread.downgrade(),
255 action_log,
256 language_registry,
257 ));
258
259 let result = cx
260 .update(|cx| {
261 tool.clone().run(
262 ToolInput::resolved(tool_input.clone()),
263 ToolCallEventStream::test().0,
264 cx,
265 )
266 })
267 .await;
268
269 let output = match result {
270 Ok(output) => output,
271 Err(output) => anyhow::bail!("Tool returned error: {}", output),
272 };
273
274 let crate::EditFileToolOutput::Success { new_text, .. } = &output else {
275 anyhow::bail!("Tool returned error output: {}", output);
276 };
277
278 if tool_input.path != eval.input_file_path {
279 anyhow::bail!(
280 "Tool path mismatch. Expected {:?}, got {:?}",
281 eval.input_file_path,
282 tool_input.path,
283 );
284 }
285
286 if new_text != &eval.expected_output_content {
287 anyhow::bail!(
288 "Output content mismatch. Expected {:?}, got {:?}",
289 eval.expected_output_content,
290 new_text,
291 );
292 }
293
294 Ok(WriteEvalOutput {
295 tool_input,
296 text_after: new_text.clone(),
297 })
298 }
299
300 async fn extract_tool_use(
301 &self,
302 request: LanguageModelRequest,
303 cx: &mut TestAppContext,
304 ) -> Result<WriteFileToolInput> {
305 let model = self.model.clone();
306 let events = cx
307 .update(|cx| {
308 let async_cx = cx.to_async();
309 cx.foreground_executor()
310 .spawn(async move { model.stream_completion(request, &async_cx).await })
311 })
312 .await
313 .map_err(|err| anyhow::anyhow!("completion error: {}", err))?;
314
315 let mut streamed_text = String::new();
316 let mut stop_reason = None;
317 let mut parse_errors = Vec::new();
318
319 let mut events = events.fuse();
320 while let Some(event) = events.next().await {
321 match event {
322 Ok(LanguageModelCompletionEvent::ToolUse(tool_use))
323 if tool_use.is_input_complete
324 && tool_use.name.as_ref() == WriteFileTool::NAME =>
325 {
326 let input: WriteFileToolInput = tool_use
327 .input
328 .parse()
329 .context("Failed to parse tool input as WriteFileToolInput")?;
330 return Ok(input);
331 }
332 Ok(LanguageModelCompletionEvent::Text(text)) => {
333 if streamed_text.len() < 2_000 {
334 streamed_text.push_str(&text);
335 }
336 }
337 Ok(LanguageModelCompletionEvent::Stop(reason)) => {
338 stop_reason = Some(reason);
339 }
340 Ok(LanguageModelCompletionEvent::ToolUseJsonParseError {
341 tool_name,
342 raw_input,
343 json_parse_error,
344 ..
345 }) if tool_name.as_ref() == WriteFileTool::NAME => {
346 parse_errors.push(format!("{json_parse_error}\nRaw input:\n{raw_input:?}"));
347 }
348 Err(err) => return Err(anyhow::anyhow!("completion error: {}", err)),
349 _ => {}
350 }
351 }
352
353 let streamed_text = streamed_text.trim();
354 let streamed_text_suffix = if streamed_text.is_empty() {
355 String::new()
356 } else {
357 format!("\nStreamed text:\n{streamed_text}")
358 };
359 let stop_reason_suffix = stop_reason
360 .map(|reason| format!("\nStop reason: {reason:?}"))
361 .unwrap_or_default();
362 let parse_errors_suffix = if parse_errors.is_empty() {
363 String::new()
364 } else {
365 format!("\nTool parse errors:\n{}", parse_errors.join("\n"))
366 };
367
368 anyhow::bail!(
369 "Stream ended without a write_file tool use{stop_reason_suffix}{parse_errors_suffix}{streamed_text_suffix}"
370 )
371 }
372}
373
374fn run_eval(eval: EvalInput) -> eval_utils::EvalOutput<()> {
375 super::run_gpui_eval(
376 |cx| {
377 async move {
378 let test = WriteToolTest::new(cx).await;
379 let result = test.eval(eval, cx).await;
380 drop(test);
381 cx.run_until_parked();
382 result
383 }
384 .boxed_local()
385 },
386 |_| eval_utils::OutcomeKind::Passed,
387 )
388}
389
390fn message(
391 role: Role,
392 content: impl IntoIterator<Item = MessageContent>,
393) -> LanguageModelRequestMessage {
394 LanguageModelRequestMessage {
395 role,
396 content: content.into_iter().collect(),
397 cache: false,
398 reasoning_details: None,
399 }
400}
401
402fn text(text: impl Into<String>) -> MessageContent {
403 MessageContent::Text(text.into())
404}
405
406fn tool_use(
407 id: impl Into<Arc<str>>,
408 name: impl Into<Arc<str>>,
409 input: impl Serialize,
410) -> MessageContent {
411 MessageContent::ToolUse(LanguageModelToolUse {
412 id: LanguageModelToolUseId::from(id.into()),
413 name: name.into(),
414 raw_input: serde_json::to_string_pretty(&input).unwrap(),
415 input: language_model::LanguageModelToolUseInput::Json(
416 serde_json::to_value(input).unwrap(),
417 ),
418 is_input_complete: true,
419 thought_signature: None,
420 })
421}
422
423fn tool_result(
424 id: impl Into<Arc<str>>,
425 name: impl Into<Arc<str>>,
426 result: impl Into<Arc<str>>,
427) -> MessageContent {
428 MessageContent::ToolResult(LanguageModelToolResult {
429 tool_use_id: LanguageModelToolUseId::from(id.into()),
430 tool_name: name.into(),
431 is_error: false,
432 content: vec![LanguageModelToolResultContent::Text(result.into())],
433 output: None,
434 })
435}
436
437async fn retry_on_rate_limit<R>(mut request: impl AsyncFnMut() -> Result<R>) -> Result<R> {
438 const MAX_RETRIES: usize = 20;
439 let mut attempt = 0;
440
441 loop {
442 attempt += 1;
443 let response = request().await;
444
445 if attempt >= MAX_RETRIES {
446 return response;
447 }
448
449 let retry_delay = match &response {
450 Ok(_) => None,
451 Err(err) => match err.downcast_ref::<LanguageModelCompletionError>() {
452 Some(err) => match &err {
453 LanguageModelCompletionError::RateLimitExceeded { retry_after, .. }
454 | LanguageModelCompletionError::ServerOverloaded { retry_after, .. } => {
455 Some(retry_after.unwrap_or(Duration::from_secs(5)))
456 }
457 LanguageModelCompletionError::UpstreamProviderError {
458 status,
459 retry_after,
460 ..
461 } => {
462 let should_retry = matches!(
463 *status,
464 StatusCode::TOO_MANY_REQUESTS | StatusCode::SERVICE_UNAVAILABLE
465 ) || status.as_u16() == 529;
466
467 if should_retry {
468 Some(retry_after.unwrap_or(Duration::from_secs(5)))
469 } else {
470 None
471 }
472 }
473 LanguageModelCompletionError::ApiReadResponseError { .. }
474 | LanguageModelCompletionError::ApiInternalServerError { .. }
475 | LanguageModelCompletionError::HttpSend { .. } => {
476 Some(Duration::from_secs(2_u64.pow((attempt - 1) as u32).min(30)))
477 }
478 _ => None,
479 },
480 _ => None,
481 },
482 };
483
484 if let Some(retry_after) = retry_delay {
485 let jitter = retry_after.mul_f64(rand::rng().random_range(0.0..1.0));
486 eprintln!("Attempt #{attempt}: Retry after {retry_after:?} + jitter of {jitter:?}");
487 #[allow(clippy::disallowed_methods)]
488 async_io::Timer::after(retry_after + jitter).await;
489 } else {
490 return response;
491 }
492 }
493}
494
495#[test]
496#[cfg_attr(not(feature = "unit-eval"), ignore)]
497fn eval_create_file() {
498 let input_file_path = "root/TODO3";
499 let expected_output_content = "todo".to_string();
500
501 eval_utils::eval(100, 1., eval_utils::NoProcessor, move || {
502 run_eval(EvalInput::new(
503 vec![
504 message(
505 User,
506 [text("Create a third todo file. Write 'todo' inside it.")],
507 ),
508 message(
509 Assistant,
510 [
511 text(indoc::formatdoc! {"
512 I'll help you create a third empty todo file.
513 First, let me examine the project structure to see if there's already a todo file, which will help me determine the appropriate name and location for the second one.
514 "}),
515 tool_use(
516 "toolu_01GAF8TtsgpjKxCr8fgQLDgR",
517 ListDirectoryTool::NAME,
518 ListDirectoryToolInput {
519 path: "root".to_string(),
520 },
521 ),
522 ],
523 ),
524 message(
525 User,
526 [tool_result(
527 "toolu_01GAF8TtsgpjKxCr8fgQLDgR",
528 ListDirectoryTool::NAME,
529 "root/TODO\nroot/TODO2\nroot/new.txt\n",
530 )],
531 ),
532 ],
533 input_file_path,
534 None,
535 expected_output_content.clone(),
536 ))
537 });
538}
539
540#[test]
541#[cfg_attr(not(feature = "unit-eval"), ignore)]
542fn eval_overwrite_file() {
543 let input_file_path = "root/notes.txt";
544 let input_file_content = "old notes\nkeep nothing\n".to_string();
545 let expected_output_content = "new notes".to_string();
546
547 eval_utils::eval(100, 1., eval_utils::NoProcessor, move || {
548 run_eval(EvalInput::new(
549 vec![message(
550 User,
551 [text(indoc::formatdoc! {"
552 Overwrite `{input_file_path}` so that its complete contents are exactly: 'new notes'
553 "})],
554 )],
555 input_file_path,
556 Some(input_file_content.clone()),
557 expected_output_content.clone(),
558 ))
559 });
560}
561