返回 CodeWhale
lib.rs
根目录 / crates / tools / src / lib.rs
1 use std::collections::HashMap;
2 use std::path::PathBuf;
3 use std::sync::Arc;
4 use std::time::Duration;
5
6 use anyhow::Result;
7 use async_trait::async_trait;
8 use codewhale_protocol::{ToolKind, ToolOutput, ToolPayload};
9 use serde::{Deserialize, Serialize};
10 use serde_json::Value;
11 use tokio::sync::{OwnedRwLockReadGuard, OwnedRwLockWriteGuard, RwLock};
12
13 mod outcome;
14 mod prepared;
15 mod resources;
16
17 pub use outcome::{ToolExecutionOutcome, ToolTerminalStatus};
18 pub use prepared::PreparedToolCall;
19 pub use resources::{ResourceClaim, schedule_non_conflicting};
20
21 tokio::task_local! {
22 static TOOL_EXECUTION_LOCK_HELD: ();
23 }
24
25 /// Capabilities that a tool may have or require.
26 #[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
27 pub enum ToolCapability {
28 /// Tool only reads data, never modifies state.
29 ReadOnly,
30 /// Tool writes to the filesystem.
31 WritesFiles,
32 /// Tool executes arbitrary shell commands.
33 ExecutesCode,
34 /// Tool makes network requests.
35 Network,
36 /// Tool can be run in a sandbox.
37 Sandboxable,
38 /// Tool requires user approval before execution.
39 RequiresApproval,
40 }
41
42 /// Approval requirement for a tool.
43 #[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
44 pub enum ApprovalRequirement {
45 /// Never needs approval: safe read-only operations.
46 #[default]
47 Auto,
48 /// Suggest approval but allow user to skip.
49 Suggest,
50 /// Always require explicit user approval.
51 Required,
52 }
53
54 /// Errors that can occur during tool execution.
55 #[derive(Debug, Clone, thiserror::Error)]
56 pub enum ToolError {
57 #[error("Failed to validate input: {message}")]
58 InvalidInput { message: String },
59 #[error("Failed to validate input: missing required field '{field}'")]
60 MissingField { field: String },
61 #[error("Failed to resolve path '{}': path escapes workspace", path.display())]
62 PathEscape { path: PathBuf },
63 #[error("Failed to execute tool: {message}")]
64 ExecutionFailed {
65 message: String,
66 /// Structured facts about the failure, in the same shape a
67 /// `ToolResult` would carry them. A process-backed tool that reports a
68 /// nonzero exit or a timeout as an error puts its `exit_code` and
69 /// `status` here, so observers (hooks, receipts) still see them.
70 metadata: Option<Value>,
71 },
72 #[error("Failed to execute tool: operation timed out after {seconds}s")]
73 Timeout { seconds: u64 },
74 #[error("Tool execution cancelled: {message}")]
75 Cancelled { message: String },
76 #[error("Failed to locate tool: {message}")]
77 NotAvailable { message: String },
78 #[error("Failed to authorize tool execution: {message}")]
79 PermissionDenied { message: String },
80 }
81
82 impl ToolError {
83 #[must_use]
84 pub fn invalid_input(msg: impl Into<String>) -> Self {
85 Self::InvalidInput {
86 message: msg.into(),
87 }
88 }
89
90 #[must_use]
91 pub fn missing_field(field: impl Into<String>) -> Self {
92 Self::MissingField {
93 field: field.into(),
94 }
95 }
96
97 #[must_use]
98 pub fn execution_failed(msg: impl Into<String>) -> Self {
99 Self::ExecutionFailed {
100 message: msg.into(),
101 metadata: None,
102 }
103 }
104
105 /// An execution failure that still carries structured result metadata,
106 /// such as a shell command's `exit_code` and `status`.
107 #[must_use]
108 pub fn execution_failed_with_metadata(msg: impl Into<String>, metadata: Value) -> Self {
109 Self::ExecutionFailed {
110 message: msg.into(),
111 metadata: Some(metadata),
112 }
113 }
114
115 /// Structured metadata the failing tool attached, if any.
116 #[must_use]
117 pub fn metadata(&self) -> Option<&Value> {
118 match self {
119 Self::ExecutionFailed { metadata, .. } => metadata.as_ref(),
120 _ => None,
121 }
122 }
123
124 #[must_use]
125 pub fn cancelled(msg: impl Into<String>) -> Self {
126 Self::Cancelled {
127 message: msg.into(),
128 }
129 }
130
131 #[must_use]
132 pub fn path_escape(path: impl Into<PathBuf>) -> Self {
133 Self::PathEscape { path: path.into() }
134 }
135
136 #[must_use]
137 pub fn not_available(msg: impl Into<String>) -> Self {
138 Self::NotAvailable {
139 message: msg.into(),
140 }
141 }
142
143 #[must_use]
144 pub fn permission_denied(msg: impl Into<String>) -> Self {
145 Self::PermissionDenied {
146 message: msg.into(),
147 }
148 }
149 }
150
151 /// Result of a tool execution.
152 #[derive(Debug, Clone, Serialize, Deserialize)]
153 pub struct ToolResult {
154 /// The output content, which may be JSON or plain text.
155 pub content: String,
156 /// Whether the execution was successful.
157 pub success: bool,
158 /// Optional structured metadata.
159 #[serde(skip_serializing_if = "Option::is_none")]
160 pub metadata: Option<Value>,
161 }
162
163 /// Provider-neutral non-text content returned alongside a tool result.
164 /// Image-producing tools return owned base64 bytes here (MCP uses its standard
165 /// `content` image blocks). A path in `ToolResult.metadata` is descriptive
166 /// metadata, never permission for the engine to read another host file.
167 /// The runtime validates format, full decode and size, retains one image per
168 /// result, and keeps omitted-image receipts with the text result.
169 #[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
170 #[serde(tag = "type", rename_all = "snake_case")]
171 pub enum ToolResultContentBlock {
172 Image { mime_type: String, data: String },
173 }
174
175 impl ToolResult {
176 /// Create a successful result with content.
177 #[must_use]
178 pub fn success(content: impl Into<String>) -> Self {
179 Self {
180 content: content.into(),
181 success: true,
182 metadata: None,
183 }
184 }
185
186 /// Create an error result with message.
187 #[must_use]
188 pub fn error(message: impl Into<String>) -> Self {
189 Self {
190 content: message.into(),
191 success: false,
192 metadata: None,
193 }
194 }
195
196 /// Create a successful result from JSON.
197 pub fn json<T: Serialize>(value: &T) -> std::result::Result<Self, serde_json::Error> {
198 Ok(Self {
199 content: serde_json::to_string(value)?,
200 success: true,
201 metadata: None,
202 })
203 }
204
205 /// Add metadata to the result.
206 #[must_use]
207 pub fn with_metadata(mut self, metadata: Value) -> Self {
208 self.metadata = Some(metadata);
209 self
210 }
211 }
212
213 /// Name the JSON type of a value the way a tool schema would spell it.
214 #[must_use]
215 pub fn json_type_name(value: &Value) -> &'static str {
216 match value {
217 Value::Null => "null",
218 Value::Bool(_) => "boolean",
219 Value::Number(_) => "number",
220 Value::String(_) => "string",
221 Value::Array(_) => "array",
222 Value::Object(_) => "object",
223 }
224 }
225
226 /// Render a value for an error message, truncated so a huge payload cannot
227 /// swamp the transcript.
228 #[must_use]
229 pub fn value_preview(value: &Value) -> String {
230 let preview = value.to_string();
231 if preview.chars().count() > 120 {
232 preview.chars().take(117).collect::<String>() + "..."
233 } else {
234 preview
235 }
236 }
237
238 /// The one error every type mismatch on a tool parameter produces.
239 ///
240 /// Names the parameter, the type that arrived, and the type the schema
241 /// declares, plus the offending value — everything the caller needs to fix
242 /// the call on the next turn without another round trip.
243 #[must_use]
244 pub fn type_mismatch(field: &str, value: &Value, expected: &str) -> ToolError {
245 ToolError::invalid_input(format!(
246 "field '{field}' must be {expected}; got {}. Received: {}",
247 json_type_name(value),
248 value_preview(value)
249 ))
250 }
251
252 /// Whether a value counts as "the caller did not supply this field".
253 ///
254 /// JSON `null` is the wire spelling of absence, so an optional field set to
255 /// `null` takes its default rather than erroring. This is the *only*
256 /// tolerance in the optional extractors, and it is uniform across all of
257 /// them: `null` means no value, and no value is exactly what a default is
258 /// for. Every other type mismatch is an error.
259 fn is_absent(value: Option<&Value>) -> bool {
260 matches!(value, None | Some(Value::Null))
261 }
262
263 /// Helper to extract a required string field from JSON input.
264 pub fn required_str<'a>(input: &'a Value, field: &str) -> std::result::Result<&'a str, ToolError> {
265 if let Some(value) = input.get(field) {
266 if let Some(string_value) = value.as_str() {
267 return Ok(string_value);
268 }
269
270 return Err(type_mismatch(field, value, "a string"));
271 }
272
273 // When the field is missing, list the fields the caller *did*
274 // supply so the model can spot the mismatch without a retry.
275 let provided: Vec<&str> = input
276 .as_object()
277 .map(|obj| obj.keys().map(|k| k.as_str()).collect())
278 .unwrap_or_default();
279 if provided.is_empty() {
280 Err(ToolError::missing_field(field))
281 } else {
282 let hint = format!(
283 "missing required field '{field}'. Input provided: {}",
284 provided.join(", ")
285 );
286 Err(ToolError::invalid_input(hint))
287 }
288 }
289
290 /// Helper to extract an optional string field from JSON input.
291 ///
292 /// A wrong type is an error, never a silent `None`. See [`type_mismatch`]
293 /// for why nothing is coerced.
294 pub fn optional_str<'a>(
295 input: &'a Value,
296 field: &str,
297 ) -> std::result::Result<Option<&'a str>, ToolError> {
298 let value = input.get(field);
299 if is_absent(value) {
300 return Ok(None);
301 }
302 let value = value.expect("is_absent covers the None case");
303 value
304 .as_str()
305 .map(Some)
306 .ok_or_else(|| type_mismatch(field, value, "a string"))
307 }
308
309 /// Read a JSON number as a non-negative integer, accepting whole-number
310 /// floats such as `200.0`.
311 ///
312 /// Several providers serialize every JSON number as a float, so a ranged
313 /// `read` arrives as `{"offset": 200.0}`. Exact integers keep the `as_u64`
314 /// fast path (so `u64::MAX` stays exact); a float is accepted only when it
315 /// is finite, non-negative, has no fractional part and fits in `u64`.
316 /// Negative, fractional, string and array values are still refused.
317 #[must_use]
318 pub fn json_nonnegative_integer(value: &Value) -> Option<u64> {
319 if let Some(number) = value.as_u64() {
320 return Some(number);
321 }
322 let float = value.as_f64()?;
323 // 2^64 is exactly representable as f64; anything at or above it overflows.
324 const U64_LIMIT: f64 = 18_446_744_073_709_551_616.0;
325 if float.is_finite() && float >= 0.0 && float.fract() == 0.0 && float < U64_LIMIT {
326 // The guards above make this cast exact.
327 return Some(float as u64);
328 }
329 None
330 }
331
332 /// Helper to extract a required u64 field from JSON input.
333 ///
334 /// Absence (field missing or `null`) is a `missing_field` error; a value
335 /// that is present but not a u64 is a [`type_mismatch`] naming the field and
336 /// the expected type, so the caller fixes the field's type instead of
337 /// re-sending it as missing.
338 pub fn required_u64(input: &Value, field: &str) -> std::result::Result<u64, ToolError> {
339 let value = input.get(field);
340 if is_absent(value) {
341 return Err(ToolError::missing_field(field));
342 }
343 let value = value.expect("is_absent covers the None case");
344 json_nonnegative_integer(value)
345 .ok_or_else(|| type_mismatch(field, value, "a non-negative integer"))
346 }
347
348 /// Helper to extract an optional u64 field with default.
349 ///
350 /// A wrong type is an error, never a silent fall back to `default`.
351 pub fn optional_u64(
352 input: &Value,
353 field: &str,
354 default: u64,
355 ) -> std::result::Result<u64, ToolError> {
356 let value = input.get(field);
357 if is_absent(value) {
358 return Ok(default);
359 }
360 let value = value.expect("is_absent covers the None case");
361 json_nonnegative_integer(value)
362 .ok_or_else(|| type_mismatch(field, value, "a non-negative integer"))
363 }
364
365 /// Helper to extract an optional bool field with default.
366 ///
367 /// A wrong type is an error, never a silent fall back to `default`. In
368 /// particular the string `"true"` is refused rather than coerced: the
369 /// default this used to fall back to is frequently the *opposite* of what
370 /// the caller asked for, and some of those defaults gate irreversible
371 /// actions.
372 pub fn optional_bool(
373 input: &Value,
374 field: &str,
375 default: bool,
376 ) -> std::result::Result<bool, ToolError> {
377 Ok(optional_bool_opt(input, field)?.unwrap_or(default))
378 }
379
380 /// Helper to extract an optional bool that has no default.
381 ///
382 /// `None` means the caller did not supply the field; a wrong type is an
383 /// error. Use this where "unset" is itself meaningful — an authority
384 /// declaration that is dropped instead of read is a restriction that
385 /// silently evaporates.
386 pub fn optional_bool_opt(
387 input: &Value,
388 field: &str,
389 ) -> std::result::Result<Option<bool>, ToolError> {
390 let value = input.get(field);
391 if is_absent(value) {
392 return Ok(None);
393 }
394 let value = value.expect("is_absent covers the None case");
395 value
396 .as_bool()
397 .map(Some)
398 .ok_or_else(|| type_mismatch(field, value, "a boolean"))
399 }
400
401 /// Descriptor that describes a tool available in the registry.
402 ///
403 /// Contains the tool's name, its JSON input/output schemas, and
404 /// execution constraints such as timeout and parallelism.
405 #[derive(Debug, Clone, Serialize, Deserialize)]
406 pub struct ToolDescriptor {
407 /// Unique name used to look up the tool.
408 pub name: String,
409 /// JSON Schema describing the tool's expected input parameters.
410 pub input_schema: Value,
411 /// JSON Schema describing the tool's output format.
412 pub output_schema: Value,
413 /// Whether multiple invocations of this tool may run concurrently.
414 pub supports_parallel_tool_calls: bool,
415 /// Optional per-call timeout in milliseconds; `None` means no timeout.
416 pub timeout_ms: Option<u64>,
417 }
418
419 /// A [`ToolDescriptor`] together with its runtime configuration.
420 ///
421 /// Wraps a `ToolDescriptor` and exposes the parallelism flag directly so the
422 /// dispatcher can check it without digging into the inner spec.
423 #[derive(Debug, Clone, Serialize, Deserialize)]
424 pub struct ConfiguredToolDescriptor {
425 /// The underlying tool descriptor.
426 pub spec: ToolDescriptor,
427 /// Whether this tool supports concurrent invocations.
428 pub supports_parallel_tool_calls: bool,
429 }
430
431 /// Identifies where a tool call originated from.
432 #[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
433 #[serde(rename_all = "snake_case")]
434 pub enum ToolCallSource {
435 /// Direct invocation from the model or user.
436 Direct,
437 /// Invocation through the JavaScript REPL environment.
438 JsRepl,
439 }
440
441 /// A tool invocation request before it has been validated and dispatched.
442 ///
443 /// Contains the tool name, its input payload, and metadata about where the
444 /// call originated.
445 #[derive(Debug, Clone, Serialize, Deserialize)]
446 pub struct ToolCall {
447 /// Name of the tool to invoke.
448 pub name: String,
449 /// The input payload for the tool.
450 pub payload: ToolPayload,
451 /// Where this call originated (direct or REPL).
452 pub source: ToolCallSource,
453 /// Optional raw tool-call identifier from the upstream provider.
454 pub raw_tool_call_id: Option<String>,
455 }
456
457 impl ToolCall {
458 /// Derive the execution subject for this call.
459 ///
460 /// For local shell payloads this returns the shell command and its
461 /// working directory; for all other payloads the tool name and the
462 /// provided `fallback_cwd` are returned instead. The third element
463 /// of the tuple is a human-readable kind label (`"shell"` or `"tool"`).
464 pub fn execution_subject(&self, fallback_cwd: &str) -> (String, String, &'static str) {
465 match &self.payload {
466 ToolPayload::LocalShell { params } => (
467 params.command.clone(),
468 params
469 .cwd
470 .clone()
471 .unwrap_or_else(|| fallback_cwd.to_string()),
472 "shell",
473 ),
474 _ => (self.name.clone(), fallback_cwd.to_string(), "tool"),
475 }
476 }
477 }
478
479 /// A validated tool invocation ready to be handled.
480 ///
481 /// Created by the registry after a [`ToolCall`] passes validation, this
482 /// carries all the context a [`ToolHandler`] needs to execute the tool.
483 #[derive(Debug, Clone)]
484 pub struct ToolInvocation {
485 /// Unique identifier for this invocation (generated or from the provider).
486 pub call_id: String,
487 /// Name of the tool being invoked.
488 pub tool_name: String,
489 /// The input payload for the tool.
490 pub payload: ToolPayload,
491 /// Where this invocation originated.
492 pub source: ToolCallSource,
493 }
494
495 /// Errors that can occur during tool dispatch and execution.
496 ///
497 /// Unlike [`ToolError`], which represents input validation failures within
498 /// a tool, `FunctionCallError` covers problems at the dispatch layer: the
499 /// tool was not found, its kind did not match, it was rejected because it
500 /// is mutating, it timed out, was cancelled, or its handler returned an
501 /// error.
502 #[derive(Debug, Clone, Serialize, Deserialize)]
503 pub enum FunctionCallError {
504 /// No tool with the given name is registered.
505 ToolNotFound { name: String },
506 /// The payload kind does not match the handler's expected kind.
507 KindMismatch { expected: ToolKind, got: ToolKind },
508 /// The tool is mutating but `allow_mutating` was `false`.
509 MutatingToolRejected { name: String },
510 /// The tool execution exceeded its configured timeout.
511 TimedOut { name: String, timeout_ms: u64 },
512 /// The tool execution was cancelled.
513 Cancelled { name: String },
514 /// The tool handler returned an error.
515 ExecutionFailed { name: String, error: String },
516 }
517
518 /// Trait implemented by concrete tool handlers.
519 ///
520 /// Each registered tool is backed by a handler that reports its kind,
521 /// whether it is mutating, and performs the actual execution.
522 #[async_trait]
523 pub trait ToolHandler: Send + Sync {
524 /// The [`ToolKind`] this handler expects (e.g. `Function` or `Mcp`).
525 fn kind(&self) -> ToolKind;
526
527 /// Returns `true` if `kind` matches this handler's expected kind.
528 ///
529 /// The default implementation compares against [`kind()`](ToolHandler::kind).
530 fn matches_kind(&self, kind: ToolKind) -> bool {
531 self.kind() == kind
532 }
533
534 /// Whether this tool performs side-effects that require user approval.
535 ///
536 /// Defaults to `false` (read-only / safe).
537 fn is_mutating(&self) -> bool {
538 false
539 }
540
541 /// Execute the tool with the given invocation context.
542 async fn handle(
543 &self,
544 invocation: ToolInvocation,
545 ) -> std::result::Result<ToolOutput, FunctionCallError>;
546 }
547
548 /// Manages concurrent tool execution via a read/write lock.
549 ///
550 /// Parallel-safe tools acquire a read lock (allowing overlap), while
551 /// serial tools acquire a write lock (exclusive access). Reentrant calls
552 /// (e.g. a tool invoking another tool) skip locking to avoid deadlock.
553 #[derive(Debug)]
554 pub struct ToolCallRuntime {
555 execution_lock: Arc<RwLock<()>>,
556 }
557
558 impl Default for ToolCallRuntime {
559 fn default() -> Self {
560 Self {
561 execution_lock: Arc::new(RwLock::new(())),
562 }
563 }
564 }
565
566 #[derive(Debug)]
567 enum ToolExecutionGuard {
568 Parallel(#[allow(dead_code)] OwnedRwLockReadGuard<()>),
569 Serial(#[allow(dead_code)] OwnedRwLockWriteGuard<()>),
570 Reentrant,
571 }
572
573 impl ToolCallRuntime {
574 async fn acquire(&self, supports_parallel: bool) -> ToolExecutionGuard {
575 if TOOL_EXECUTION_LOCK_HELD.try_with(|_| ()).is_ok() {
576 return ToolExecutionGuard::Reentrant;
577 }
578
579 if supports_parallel {
580 ToolExecutionGuard::Parallel(self.execution_lock.clone().read_owned().await)
581 } else {
582 ToolExecutionGuard::Serial(self.execution_lock.clone().write_owned().await)
583 }
584 }
585 }
586
587 /// Central registry that maps tool names to their specs and handlers.
588 ///
589 /// Use [`register()`](ToolRegistry::register) to add tools, then
590 /// [`dispatch()`](ToolRegistry::dispatch) to invoke them. The registry
591 /// owns a [`ToolCallRuntime`] that manages concurrent execution.
592 #[derive(Default)]
593 pub struct ToolRegistry {
594 handlers: HashMap<String, Arc<dyn ToolHandler>>,
595 specs: HashMap<String, ConfiguredToolDescriptor>,
596 runtime: ToolCallRuntime,
597 }
598
599 impl ToolRegistry {
600 /// Register a tool with its specification and handler.
601 ///
602 /// The tool's name is taken from `spec.name`. Returns an error if
603 /// registration fails (currently infallible, but the `Result` is
604 /// reserved for future validation).
605 pub fn register(&mut self, spec: ToolDescriptor, handler: Arc<dyn ToolHandler>) -> Result<()> {
606 let name = spec.name.clone();
607 self.specs.insert(
608 name.clone(),
609 ConfiguredToolDescriptor {
610 supports_parallel_tool_calls: spec.supports_parallel_tool_calls,
611 spec,
612 },
613 );
614 self.handlers.insert(name, handler);
615 Ok(())
616 }
617
618 /// Validate and execute a tool call.
619 ///
620 /// Looks up the tool by name, verifies the payload kind matches the
621 /// handler, enforces the `allow_mutating` guard, acquires the
622 /// appropriate execution lock, and forwards the call to the handler.
623 /// Returns a [`FunctionCallError`] if any validation step fails or
624 /// the handler returns an error.
625 pub async fn dispatch(
626 &self,
627 call: ToolCall,
628 allow_mutating: bool,
629 ) -> std::result::Result<ToolOutput, FunctionCallError> {
630 let handler = self.handlers.get(&call.name).cloned().ok_or_else(|| {
631 FunctionCallError::ToolNotFound {
632 name: call.name.clone(),
633 }
634 })?;
635 let configured =
636 self.specs
637 .get(&call.name)
638 .cloned()
639 .ok_or_else(|| FunctionCallError::ToolNotFound {
640 name: call.name.clone(),
641 })?;
642
643 let payload_kind = tool_payload_kind(&call.payload);
644 let expected = handler.kind();
645 if !handler.matches_kind(payload_kind) {
646 return Err(FunctionCallError::KindMismatch {
647 expected,
648 got: payload_kind,
649 });
650 }
651 if handler.is_mutating() && !allow_mutating {
652 return Err(FunctionCallError::MutatingToolRejected { name: call.name });
653 }
654
655 let invocation = ToolInvocation {
656 call_id: call
657 .raw_tool_call_id
658 .clone()
659 .unwrap_or_else(|| format!("tool-call-{}", uuid::Uuid::new_v4())),
660 tool_name: call.name.clone(),
661 payload: call.payload,
662 source: call.source,
663 };
664
665 let _guard = self
666 .runtime
667 .acquire(configured.supports_parallel_tool_calls)
668 .await;
669
670 TOOL_EXECUTION_LOCK_HELD
671 .scope(
672 (),
673 self.execute_with_timeout(handler, configured.spec.timeout_ms, invocation),
674 )
675 .await
676 }
677
678 async fn execute_with_timeout(
679 &self,
680 handler: Arc<dyn ToolHandler>,
681 timeout_ms: Option<u64>,
682 invocation: ToolInvocation,
683 ) -> std::result::Result<ToolOutput, FunctionCallError> {
684 if let Some(timeout_ms) = timeout_ms {
685 let name = invocation.tool_name.clone();
686 match tokio::time::timeout(
687 Duration::from_millis(timeout_ms),
688 handler.handle(invocation),
689 )
690 .await
691 {
692 Ok(result) => result,
693 Err(_) => Err(FunctionCallError::TimedOut { name, timeout_ms }),
694 }
695 } else {
696 handler.handle(invocation).await
697 }
698 }
699 }
700
701 fn tool_payload_kind(payload: &ToolPayload) -> ToolKind {
702 match payload {
703 ToolPayload::Mcp { .. } => ToolKind::Mcp,
704 ToolPayload::Function { .. }
705 | ToolPayload::Custom { .. }
706 | ToolPayload::LocalShell { .. } => ToolKind::Function,
707 }
708 }
709
710 #[cfg(test)]
711 mod tests {
712 use serde_json::json;
713
714 use super::*;
715
716 #[test]
717 fn tool_result_success_sets_plain_content() {
718 let content = "operation completed successfully";
719 let result = ToolResult::success(content);
720
721 assert!(result.success);
722 assert_eq!(result.content, content);
723 assert!(result.metadata.is_none());
724 }
725
726 #[test]
727 fn tool_result_json_round_trips_content() {
728 let result = ToolResult::json(&json!({"ok": true})).expect("json");
729 assert!(result.success);
730 let content: serde_json::Value =
731 serde_json::from_str(&result.content).expect("content is valid json");
732 assert_eq!(content, json!({"ok": true}));
733 }
734
735 #[test]
736 fn helper_extractors_validate_shape() {
737 let input = json!({"name": "demo", "count": 7, "enabled": true});
738 assert_eq!(required_str(&input, "name").expect("name"), "demo");
739 assert_eq!(optional_str(&input, "name").unwrap(), Some("demo"));
740 assert_eq!(optional_str(&input, "missing").unwrap(), None);
741 assert_eq!(optional_str(&json!({"name": null}), "name").unwrap(), None);
742 assert_eq!(optional_u64(&input, "count", 0).unwrap(), 7);
743 assert!(optional_bool(&input, "enabled", false).unwrap());
744 // "name" is present but a string: a type mismatch, not a missing
745 // field, so the caller fixes the type instead of re-sending the name.
746 let err = required_u64(&input, "name")
747 .expect_err("a present string is not a missing u64")
748 .to_string();
749 assert!(
750 err.contains("field 'name' must be a non-negative integer"),
751 "{err}"
752 );
753 }
754
755 /// The rule, stated once: an optional parameter of the wrong JSON type is
756 /// an error that names the parameter, what arrived, and what was wanted.
757 /// `null` alone means "absent" and takes the default.
758 #[test]
759 fn optional_extractors_refuse_type_mismatches_instead_of_defaulting() {
760 // The shipping bug: a stringy "true" became the default `false`,
761 // which for `dry_run` is the opposite of what the caller asked and
762 // gates an irreversible action.
763 let err = optional_bool(&json!({"dry_run": "true"}), "dry_run", false)
764 .expect_err("a stringy bool must not become the default")
765 .to_string();
766 assert!(err.contains("dry_run"), "{err}");
767 assert!(err.contains("must be a boolean"), "{err}");
768 assert!(err.contains("got string"), "{err}");
769 assert!(err.contains("\"true\""), "{err}");
770
771 for bad in [json!("true"), json!(1), json!(0), json!([]), json!({})] {
772 assert!(
773 optional_bool(&json!({"flag": bad}), "flag", false).is_err(),
774 "optional_bool accepted {bad}"
775 );
776 }
777 for bad in [json!("7"), json!(-1), json!(1.5), json!(true), json!([7])] {
778 assert!(
779 optional_u64(&json!({"n": bad}), "n", 42).is_err(),
780 "optional_u64 accepted {bad}"
781 );
782 }
783 for bad in [json!(7), json!(true), json!(["a"]), json!({"a": 1})] {
784 assert!(
785 optional_str(&json!({"s": bad}), "s").is_err(),
786 "optional_str accepted {bad}"
787 );
788 }
789
790 // `null` is the wire spelling of absence, uniformly across all three.
791 assert!(optional_bool(&json!({"flag": null}), "flag", true).unwrap());
792 assert_eq!(optional_u64(&json!({"n": null}), "n", 42).unwrap(), 42);
793 assert_eq!(optional_str(&json!({"s": null}), "s").unwrap(), None);
794 }
795
796 #[test]
797 fn type_mismatch_truncates_a_huge_offending_value() {
798 let big = Value::String("x".repeat(500));
799 let err = type_mismatch("body", &big, "a boolean").to_string();
800 assert!(err.contains("body"), "{err}");
801 assert!(err.ends_with("..."), "{err}");
802 assert!(err.chars().count() < 250, "{err}");
803 }
804
805 #[test]
806 fn required_u64_distinguishes_missing_from_type_mismatch() {
807 // Absent (or null) is a missing-field error.
808 assert!(matches!(
809 required_u64(&json!({}), "count"),
810 Err(ToolError::MissingField { .. })
811 ));
812 assert!(matches!(
813 required_u64(&json!({"count": null}), "count"),
814 Err(ToolError::MissingField { .. })
815 ));
816
817 // Present and valid values pass through, including the extremes.
818 assert_eq!(required_u64(&json!({"count": 42}), "count").unwrap(), 42);
819 assert_eq!(
820 required_u64(&json!({"count": u64::MAX}), "count").unwrap(),
821 u64::MAX
822 );
823
824 // Whole-number floats are integers some providers send as `200.0`.
825 assert_eq!(
826 required_u64(&json!({"count": 200.0}), "count").unwrap(),
827 200
828 );
829 assert_eq!(required_u64(&json!({"count": 0.0}), "count").unwrap(), 0);
830 assert_eq!(
831 optional_u64(&json!({"count": 200.0}), "count", 7).unwrap(),
832 200
833 );
834
835 // Present but wrongly typed is a type mismatch naming the field and
836 // the expected type — never a missing-field misdirection. Floats that
837 // are negative, fractional or beyond u64 stay refused.
838 for value in [
839 json!(-1),
840 json!(-1.0),
841 json!(2.5),
842 json!(1.8446744073709552e19),
843 json!("42"),
844 ] {
845 let err = required_u64(&json!({"count": value}), "count")
846 .expect_err("wrong type must not look missing")
847 .to_string();
848 assert!(
849 err.contains("field 'count' must be a non-negative integer"),
850 "{err}"
851 );
852 }
853 }
854
855 #[test]
856 fn required_str_reports_provided_fields_on_missing_required_field() {
857 let input = json!({"path": "src/lib.rs", "content": "new body"});
858 let err = required_str(&input, "replace").expect_err("replace is missing");
859 let message = err.to_string();
860 assert!(message.contains("missing required field 'replace'"));
861 assert!(message.contains("Input provided:"));
862 assert!(message.contains("path"));
863 assert!(message.contains("content"));
864 }
865
866 #[test]
867 fn required_str_reports_wrong_type_when_field_exists() {
868 let input = json!({"replace": [{"path": "src/lib.rs", "content": "new body"}]});
869 let err = required_str(&input, "replace").expect_err("replace has wrong type");
870 let message = err.to_string();
871 assert!(message.contains("field 'replace' must be a string"));
872 assert!(message.contains("got array"));
873 assert!(message.contains(r#""content":"new body""#));
874 assert!(message.contains(r#""path":"src/lib.rs""#));
875 }
876
877 #[test]
878 fn tool_error_display_matches_legacy_text() {
879 let err = ToolError::missing_field("path");
880 assert_eq!(
881 err.to_string(),
882 "Failed to validate input: missing required field 'path'"
883 );
884 }
885
886 #[test]
887 fn tool_error_missing_field_constructor() {
888 let err = ToolError::missing_field("my_field");
889 assert!(matches!(err, ToolError::MissingField { field } if field == "my_field"));
890 }
891
892 #[test]
893 fn tool_error_not_available_displays_reason() {
894 let err = ToolError::not_available("custom tool not found");
895
896 assert!(matches!(err, ToolError::NotAvailable { .. }));
897 assert_eq!(
898 err.to_string(),
899 "Failed to locate tool: custom tool not found"
900 );
901 }
902
903 #[test]
904 fn tool_error_permission_denied_displays_reason() {
905 let err = ToolError::permission_denied("unauthorized user");
906
907 assert!(matches!(err, ToolError::PermissionDenied { .. }));
908 assert_eq!(
909 err.to_string(),
910 "Failed to authorize tool execution: unauthorized user"
911 );
912 }
913
914 #[test]
915 fn tool_error_execution_failed_displays_reason() {
916 let err = ToolError::execution_failed("process crashed");
917
918 assert!(
919 matches!(err, ToolError::ExecutionFailed { ref message, .. } if message == "process crashed")
920 );
921 assert_eq!(err.to_string(), "Failed to execute tool: process crashed");
922 }
923
924 #[test]
925 fn tool_error_invalid_input_creates_correct_variant() {
926 let err = ToolError::invalid_input("test invalid message");
927 match err {
928 ToolError::InvalidInput { message } => {
929 assert_eq!(message, "test invalid message");
930 }
931 _ => panic!("Expected ToolError::InvalidInput, got {err:?}"),
932 }
933 }
934
935 #[test]
936 fn tool_error_path_escape_display() {
937 let path = std::path::PathBuf::from("../outside");
938 let err = ToolError::path_escape(path);
939 assert_eq!(
940 err.to_string(),
941 "Failed to resolve path '../outside': path escapes workspace"
942 );
943 }
944
945 #[test]
946 fn tool_call_execution_subject_uses_local_shell_command_and_cwd() {
947 let call = ToolCall {
948 name: "shell".to_string(),
949 payload: ToolPayload::LocalShell {
950 params: codewhale_protocol::LocalShellParams {
951 command: "ls -l".to_string(),
952 cwd: Some("/custom/dir".to_string()),
953 timeout_ms: None,
954 },
955 },
956 source: ToolCallSource::Direct,
957 raw_tool_call_id: None,
958 };
959
960 assert_eq!(
961 call.execution_subject("/fallback/dir"),
962 ("ls -l".to_string(), "/custom/dir".to_string(), "shell")
963 );
964 }
965
966 #[test]
967 fn tool_call_execution_subject_falls_back_for_shell_without_cwd() {
968 let call = ToolCall {
969 name: "shell".to_string(),
970 payload: ToolPayload::LocalShell {
971 params: codewhale_protocol::LocalShellParams {
972 command: "echo hello".to_string(),
973 cwd: None,
974 timeout_ms: None,
975 },
976 },
977 source: ToolCallSource::Direct,
978 raw_tool_call_id: None,
979 };
980
981 assert_eq!(
982 call.execution_subject("/fallback/dir"),
983 (
984 "echo hello".to_string(),
985 "/fallback/dir".to_string(),
986 "shell"
987 )
988 );
989 }
990
991 #[test]
992 fn tool_call_execution_subject_uses_tool_name_for_non_shell_payloads() {
993 let call = ToolCall {
994 name: "my_tool".to_string(),
995 payload: ToolPayload::Function {
996 arguments: "{}".to_string(),
997 },
998 source: ToolCallSource::Direct,
999 raw_tool_call_id: None,
1000 };
1001
1002 assert_eq!(
1003 call.execution_subject("/fallback/dir"),
1004 ("my_tool".to_string(), "/fallback/dir".to_string(), "tool")
1005 );
1006 }
1007 }
1008
1008 lines RUST