1use std::path::{Path, PathBuf};
12
13use async_trait::async_trait;
14
15use mermaid_domain::{ToolDefinition, ToolMetadata, ToolOutcome, ToolRunMetadata};
16use mermaid_model::constants::MAX_PATCH_FILE_BYTES;
17use mermaid_model::diff::{DisplayDiff, MAX_DISPLAY_DIFF_LINES, generate_display_diff};
18use mermaid_runtime::apply_patch::{Hunk, UpdateFileChunk, derive_new_contents, parse_patch};
21
22use super::super::ctx::ExecContext;
23use super::ToolExecutor;
24use super::filesystem::{MutationGate, after_file_mutation, diff_summary, mutation_policy_outcome};
25use super::path_safety::{AllowedRoots, PathContainment, ResolvedInRoot, resolve_in_roots};
26
27const APPLY_PATCH_DESCRIPTION: &str = "Edit files with a patch. Pass `patch` as one string in this exact envelope:\n*** Begin Patch\n*** Update File: <path>\n@@ <optional anchor line, e.g. a function signature>\n <unchanged context line>\n-<line to remove>\n+<line to add>\n*** End Patch\nUse '*** Add File: <path>' then '+'-prefixed lines to create a file; '*** Delete File: <path>' to remove one; '*** Move to: <path>' immediately after an Update File line to rename. Include a few unchanged context lines (prefixed with a space) around each change so the edit can be located; matching tolerates whitespace/quote drift. Paths may be relative to the project directory or absolute.";
28
29pub struct ApplyPatchTool;
31
32#[async_trait]
33impl ToolExecutor for ApplyPatchTool {
34 fn name(&self) -> &'static str {
35 "apply_patch"
36 }
37
38 fn schema(&self) -> ToolDefinition {
39 ToolDefinition {
40 name: "apply_patch".to_string(),
41 description: APPLY_PATCH_DESCRIPTION.to_string(),
42 input_schema: serde_json::json!({
43 "type": "object",
44 "properties": {
45 "patch": {
46 "type": "string",
47 "description": "The full patch envelope, from '*** Begin Patch' to '*** End Patch'."
48 }
49 },
50 "required": ["patch"]
51 }),
52 }
53 }
54
55 async fn execute(&self, args: serde_json::Value, ctx: ExecContext) -> ToolOutcome {
56 let start = std::time::Instant::now();
57 let Some(patch) = args.get("patch").and_then(|v| v.as_str()) else {
58 return ToolOutcome::error("apply_patch requires 'patch' (string)", 0.0);
59 };
60 let hunks = match parse_patch(patch) {
61 Ok(h) => h,
62 Err(e) => return ToolOutcome::error(format!("apply_patch: {e}"), 0.0),
63 };
64 let (ops, paths) = match plan_ops(&ctx, &hunks) {
65 Ok(v) => v,
66 Err(e) => return ToolOutcome::error(format!("apply_patch: {e}"), 0.0),
67 };
68
69 let summary_path = ops
70 .first()
71 .map(PlannedOp::display)
72 .unwrap_or_default()
73 .to_string();
74 let pending_action = serde_json::json!({
75 "tool": "apply_patch",
76 "args": { "patch": patch },
77 "workdir": ctx.workdir.display().to_string(),
78 "turn_id": ctx.turn.0,
79 "call_id": ctx.call_id.0,
80 "task_id": ctx.task_id.clone(),
81 });
82 let plan_write = match mutation_policy_outcome(
86 &ctx,
87 "apply_patch",
88 &summary_path,
89 &paths.project,
90 pending_action,
91 paths.containment,
92 )
93 .await
94 {
95 MutationGate::Blocked(outcome) => return *outcome,
96 MutationGate::Proceed { plan_write } => plan_write,
97 };
98
99 let _guards = tokio::select! {
102 biased;
103 _ = ctx.token.cancelled() => return ToolOutcome::cancelled(),
104 g = super::path_lock::lock_paths(&paths.all) => g,
105 };
106 if ctx.config.safety.checkpoint_on_mutation
109 && !paths.project.is_empty()
110 && let Err(e) = mermaid_runtime::create_checkpoint_for_task(
111 &ctx.workdir,
112 &paths.project,
113 Some(serde_json::json!({ "tool": "apply_patch" })),
114 ctx.checkpoint_origin(),
115 )
116 {
117 return ToolOutcome::error(format!("apply_patch checkpoint failed: {e}"), 0.0);
118 }
119
120 let report = tokio::select! {
121 biased;
122 _ = ctx.token.cancelled() => return ToolOutcome::cancelled(),
123 result = tokio::task::spawn_blocking(move || apply_all_blocking(&ops)) => {
124 match result {
125 Ok(Ok(report)) => report,
126 Ok(Err(e)) => {
127 return ToolOutcome::error(format!("apply_patch: {e}"), start.elapsed().as_secs_f64());
128 },
129 Err(e) => {
130 return ToolOutcome::error(format!("apply_patch join error: {e}"), start.elapsed().as_secs_f64());
131 },
132 }
133 }
134 };
135 after_file_mutation(&ctx, "apply_patch", &summary_path);
136 build_outcome(report, start.elapsed().as_secs_f64(), plan_write)
137 }
138}
139
140enum PlannedOp {
144 Add {
145 root: PathBuf,
146 rel: PathBuf,
147 display: String,
148 contents: String,
149 },
150 Delete {
151 root: PathBuf,
152 rel: PathBuf,
153 display: String,
154 },
155 Update {
156 src_root: PathBuf,
157 src_rel: PathBuf,
158 dst_root: PathBuf,
159 dst_rel: PathBuf,
160 src_display: String,
161 dst_display: String,
162 chunks: Vec<UpdateFileChunk>,
163 },
164}
165
166impl PlannedOp {
167 fn display(&self) -> &str {
168 match self {
169 Self::Add { display, .. } | Self::Delete { display, .. } => display,
170 Self::Update { dst_display, .. } => dst_display,
171 }
172 }
173}
174
175struct PatchPaths {
180 all: Vec<PathBuf>,
181 project: Vec<PathBuf>,
182 containment: PathContainment,
186}
187
188fn plan_ops(ctx: &ExecContext, hunks: &[Hunk]) -> Result<(Vec<PlannedOp>, PatchPaths), String> {
192 let roots = AllowedRoots::new(&ctx.workdir, ctx.scratchpad.as_deref());
193 let mut ops = Vec::new();
194 let mut all: Vec<PathBuf> = Vec::new();
195 let mut project: Vec<PathBuf> = Vec::new();
196 let mut any_external = false;
197 let mut remember = |r: &ResolvedInRoot, all: &mut Vec<PathBuf>, project: &mut Vec<PathBuf>| {
198 if !all.contains(&r.abs) {
199 all.push(r.abs.clone());
200 }
201 if r.containment != PathContainment::Scratchpad && !project.contains(&r.abs) {
202 project.push(r.abs.clone());
203 }
204 any_external |= r.containment == PathContainment::External;
205 };
206 for hunk in hunks {
207 match hunk {
208 Hunk::AddFile { path, contents } => {
209 let raw = path.to_string_lossy().to_string();
210 let resolved = resolve_in_roots(&roots, &raw)?;
211 remember(&resolved, &mut all, &mut project);
212 ops.push(PlannedOp::Add {
213 root: resolved.root,
214 rel: resolved.rel,
215 display: raw,
216 contents: contents.clone(),
217 });
218 },
219 Hunk::DeleteFile { path } => {
220 let raw = path.to_string_lossy().to_string();
221 let resolved = resolve_in_roots(&roots, &raw)?;
222 remember(&resolved, &mut all, &mut project);
223 ops.push(PlannedOp::Delete {
224 root: resolved.root,
225 rel: resolved.rel,
226 display: raw,
227 });
228 },
229 Hunk::UpdateFile {
230 path,
231 move_path,
232 chunks,
233 } => {
234 let raw = path.to_string_lossy().to_string();
235 let src = resolve_in_roots(&roots, &raw)?;
236 remember(&src, &mut all, &mut project);
237 let (dst_root, dst_rel, dst_display) = match move_path {
238 Some(mv) => {
239 let mv_raw = mv.to_string_lossy().to_string();
240 let dst = resolve_in_roots(&roots, &mv_raw)?;
241 remember(&dst, &mut all, &mut project);
242 (dst.root, dst.rel, mv_raw)
243 },
244 None => (src.root.clone(), src.rel.clone(), raw.clone()),
245 };
246 ops.push(PlannedOp::Update {
247 src_root: src.root,
248 src_rel: src.rel,
249 dst_root,
250 dst_rel,
251 src_display: raw,
252 dst_display,
253 chunks: chunks.clone(),
254 });
255 },
256 }
257 }
258 all.sort();
259 project.sort();
260 let containment = if any_external {
261 PathContainment::External
262 } else if !all.is_empty() && project.is_empty() {
263 PathContainment::Scratchpad
264 } else {
265 PathContainment::Project
266 };
267 Ok((
268 ops,
269 PatchPaths {
270 all,
271 project,
272 containment,
273 },
274 ))
275}
276
277fn apply_all_blocking(ops: &[PlannedOp]) -> Result<ApplyReport, String> {
280 let mut report = ApplyReport::default();
281 for op in ops {
282 match op {
283 PlannedOp::Add {
284 root,
285 rel,
286 display,
287 contents,
288 } => {
289 let body = if contents.is_empty() || contents.ends_with('\n') {
292 contents.clone()
293 } else {
294 format!("{contents}\n")
295 };
296 ensure_parent(root, rel)?;
297 mermaid_runtime::write_atomic_beneath(root, rel, body.as_bytes())
298 .map_err(|e| format!("{display}: {e}"))?;
299 report.added.push(display.clone());
300 report.push_diff(&format!("A {display}"), &generate_display_diff("", &body));
301 },
302 PlannedOp::Delete { root, rel, display } => {
303 mermaid_runtime::remove_file_beneath(root, rel)
304 .map_err(|e| format!("{display}: {e}"))?;
305 report.deleted.push(display.clone());
306 report.push_line(format!("=== D {display} ==="));
307 },
308 PlannedOp::Update {
309 src_root,
310 src_rel,
311 dst_root,
312 dst_rel,
313 src_display,
314 dst_display,
315 chunks,
316 } => {
317 let original = read_capped_beneath(src_root, src_rel, MAX_PATCH_FILE_BYTES)
318 .map_err(|e| format!("{src_display}: {e}"))?
319 .ok_or_else(|| {
320 format!(
321 "{src_display}: file too large to patch safely (> {MAX_PATCH_FILE_BYTES} bytes)"
322 )
323 })?;
324 let applied = derive_new_contents(&original, chunks)
325 .map_err(|e| format!("{dst_display}: {e}"))?;
326 report.fuzzy |= applied.fuzzy;
327 ensure_parent(dst_root, dst_rel)?;
328 mermaid_runtime::write_atomic_beneath(
329 dst_root,
330 dst_rel,
331 applied.new_contents.as_bytes(),
332 )
333 .map_err(|e| format!("{dst_display}: {e}"))?;
334 let diff = generate_display_diff(&original, &applied.new_contents);
335 if src_root == dst_root && src_rel == dst_rel {
336 report.modified.push(dst_display.clone());
337 report.push_diff(&format!("M {dst_display}"), &diff);
338 } else {
339 mermaid_runtime::remove_file_beneath(src_root, src_rel)
340 .map_err(|e| format!("{src_display}: {e}"))?;
341 report
342 .renamed
343 .push((src_display.clone(), dst_display.clone()));
344 report.push_diff(&format!("R {src_display} -> {dst_display}"), &diff);
345 }
346 },
347 }
348 }
349 Ok(report)
350}
351
352fn ensure_parent(root: &Path, rel: &Path) -> Result<(), String> {
353 if let Some(parent) = rel.parent()
354 && !parent.as_os_str().is_empty()
355 {
356 mermaid_runtime::create_dir_all_beneath(root, parent)
357 .map_err(|e| format!("{}: {e}", rel.display()))?;
358 }
359 Ok(())
360}
361
362fn read_capped_beneath(root: &Path, rel: &Path, cap: usize) -> std::io::Result<Option<String>> {
366 use std::io::Read;
367 let file = mermaid_runtime::open_beneath(root, rel, mermaid_runtime::OpenIntent::Read)?;
368 let mut buf = Vec::new();
369 file.take(cap as u64 + 1).read_to_end(&mut buf)?;
370 if buf.len() > cap {
371 return Ok(None);
372 }
373 Ok(Some(String::from_utf8_lossy(&buf).into_owned()))
374}
375
376#[derive(Default)]
379struct ApplyReport {
380 added: Vec<String>,
381 modified: Vec<String>,
382 deleted: Vec<String>,
383 renamed: Vec<(String, String)>,
384 fuzzy: bool,
385 added_lines: usize,
386 removed_lines: usize,
387 diff_lines: Vec<String>,
388 diff_truncated: bool,
389}
390
391impl ApplyReport {
392 fn push_line(&mut self, line: String) {
393 if self.diff_lines.len() < MAX_DISPLAY_DIFF_LINES {
394 self.diff_lines.push(line);
395 } else {
396 self.diff_truncated = true;
397 }
398 }
399
400 fn push_diff(&mut self, header: &str, diff: &DisplayDiff) {
401 self.added_lines += diff.added;
402 self.removed_lines += diff.removed;
403 self.push_line(format!("=== {header} ==="));
404 for line in diff.display_diff.lines() {
405 self.push_line(line.to_string());
406 }
407 self.diff_truncated |= diff.truncated;
408 }
409}
410
411fn build_outcome(report: ApplyReport, duration_secs: f64, plan_write: bool) -> ToolOutcome {
412 let total =
413 report.added.len() + report.modified.len() + report.deleted.len() + report.renamed.len();
414 let mut lines = vec![format!("Applied patch: {total} file(s)")];
415 lines.extend(report.added.iter().map(|p| format!("A {p}")));
416 lines.extend(report.modified.iter().map(|p| format!("M {p}")));
417 lines.extend(report.renamed.iter().map(|(a, b)| format!("R {a} -> {b}")));
418 lines.extend(report.deleted.iter().map(|p| format!("D {p}")));
419 if report.fuzzy {
420 lines.push(
421 "note: one or more hunks matched with fuzzy (whitespace/Unicode) context; verify the result."
422 .to_string(),
423 );
424 }
425 let model_content = lines.join("\n");
426 ToolOutcome::success(
427 model_content,
428 diff_summary(report.added_lines, report.removed_lines, duration_secs),
429 duration_secs,
430 )
431 .with_metadata(ToolRunMetadata {
432 detail: ToolMetadata::ApplyPatch {
433 added: report.added,
434 modified: report.modified,
435 deleted: report.deleted,
436 renamed: report.renamed,
437 fuzzy: report.fuzzy,
438 },
439 display_diff: Some(report.diff_lines.join("\n")),
440 diff_truncated: report.diff_truncated,
441 lines_added: report.added_lines,
442 lines_removed: report.removed_lines,
443 plan_file_written: plan_write,
444 ..ToolRunMetadata::default()
445 })
446}
447
448#[cfg(test)]
449mod tests;