Skip to repository content

tenant.openagents/omega

No repository description is available.

OpenAgents Git authority 2026-07-28T07:47:25.062Z Public web read
NIP-34 coordinate30617:7649603503856e5148d571eac2766b288a8ff1e9e35d380337a1d2b0015b4f92:omega
MaintainersHidden in public view
References2 branches · 1 tag
Read-only clonegit clone https://openagents.com/git/tenant.openagents/omega.git
Browse files

write_file.rs

561 lines · 19.4 KB · rust
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
Served at tenant.openagents/omega Member data and write actions are omitted.