1use std::collections::HashMap;
2use std::sync::Arc;
3
4use async_trait::async_trait;
5use serde_json::{Value, json};
6use tokio::sync::mpsc;
7
8use crate::engine::SessionStore;
9use crate::types::{AgentError, AgentResult, RuntimeEvent, SessionId, UserEvent};
10
11pub mod auto_continue;
12pub mod policy;
13pub mod update_plan;
14
15pub use auto_continue::AutoContinueTool;
16pub use update_plan::UpdatePlanTool;
17
18pub use policy::{DenyAllToolPolicy, ToolDecision, ToolPolicy};
19
20pub use agent_types::{
22 ActivationContext, Content, ToolExposure, ToolMetadata, content_details, content_text,
23};
24
25#[derive(Clone)]
26pub struct ToolContext {
27 pub session_id: SessionId,
28 pub user_event_tx: mpsc::UnboundedSender<UserEvent>,
31 pub llm_client: Option<Arc<dyn llm_trait::LlmProvider>>,
32 pub session_store: Option<Arc<dyn SessionStore>>,
33 pub language: crate::types::Language,
36 pub cancel_token: tokio_util::sync::CancellationToken,
38 pub max_output_chars: Option<usize>,
43 pub(crate) event_bus: crate::engine::EventBus,
46}
47
48impl ToolContext {
49 pub fn emit_user_event(&self, event: UserEvent) {
57 let _ = self.user_event_tx.send(event.clone());
58 self.event_bus.emit(RuntimeEvent::UserEvent {
59 session_id: self.session_id.clone(),
60 event,
61 agent_id: None,
62 trace_id: None,
63 });
64 }
65
66 pub fn emit_progress(&self, text: impl Into<String>) {
68 self.emit_user_event(UserEvent::Progress { text: text.into() });
69 }
70
71 pub fn emit_partial_result(
74 &self,
75 tool_call_id: &str,
76 content: impl Into<String>,
77 is_partial: bool,
78 ) {
79 self.emit_user_event(UserEvent::ToolPartialResult {
80 tool_call_id: tool_call_id.to_string(),
81 content: content.into(),
82 is_partial,
83 });
84 }
85
86 pub fn for_test() -> Self {
90 let (tx, _rx) = mpsc::unbounded_channel();
91 ToolContext {
92 session_id: SessionId::new(0),
93 user_event_tx: tx,
94 llm_client: None,
95 session_store: None,
96 language: crate::types::Language::En,
97 cancel_token: tokio_util::sync::CancellationToken::new(),
98 max_output_chars: None,
99 event_bus: crate::engine::EventBus::new(1),
100 }
101 }
102}
103
104#[async_trait]
105pub trait Tool: Send + Sync {
106 fn name(&self) -> &'static str;
107 fn description(&self) -> &'static str;
109 fn schema(&self) -> Value;
112 async fn call(&self, args: &Value, ctx: &ToolContext) -> AgentResult<Vec<Content>>;
113
114 fn timeout_ms(&self) -> Option<u64> {
120 None }
122
123 fn metadata(&self) -> ToolMetadata {
130 ToolMetadata {
131 name: self.name().to_string(),
132 description: self.description().to_string(),
133 origin: "custom".to_string(),
134 version: "unknown".to_string(),
135 requirements: vec![],
136 }
137 }
138
139 fn exposure(&self) -> ToolExposure {
145 ToolExposure::Direct
146 }
147
148 fn should_activate(&self, _ctx: &ActivationContext) -> bool {
154 true
155 }
156}
157
158#[async_trait]
159pub trait TypedTool: Send + Sync {
160 type Args: serde::de::DeserializeOwned + schemars::JsonSchema;
161 type Output: serde::Serialize;
162
163 fn name(&self) -> &'static str;
164 fn description(&self) -> &'static str;
165 async fn call_typed(&self, args: Self::Args, ctx: &ToolContext) -> AgentResult<Self::Output>;
166
167 fn format_output(&self, output: Self::Output) -> Content {
168 match serde_json::to_value(&output) {
173 Ok(serde_json::Value::String(s)) => Content::text(s),
174 Ok(other) => Content::text(other.to_string()),
175 Err(_) => Content::text(String::new()),
176 }
177 }
178
179 fn origin(&self) -> &'static str {
181 "custom"
182 }
183
184 fn version(&self) -> &'static str {
186 "unknown"
187 }
188
189 fn exposure(&self) -> ToolExposure {
191 ToolExposure::Direct
192 }
193
194 fn should_activate(&self, _ctx: &ActivationContext) -> bool {
196 true
197 }
198}
199
200#[async_trait]
201impl<T: TypedTool + Send + Sync + 'static> Tool for T {
202 fn name(&self) -> &'static str {
203 TypedTool::name(self)
204 }
205
206 fn description(&self) -> &'static str {
207 TypedTool::description(self)
208 }
209
210 fn schema(&self) -> Value {
211 let settings = schemars::generate::SchemaSettings::draft07().with(|s| {
217 s.inline_subschemas = true;
218 s.meta_schema = None;
219 });
220 let generator = schemars::SchemaGenerator::new(settings);
221 let schema = generator.into_root_schema_for::<T::Args>();
222 serde_json::to_value(schema).unwrap_or(Value::Null)
223 }
224
225 fn metadata(&self) -> ToolMetadata {
226 ToolMetadata {
227 name: self.name().to_string(),
228 description: self.description().to_string(),
229 origin: self.origin().to_string(),
230 version: self.version().to_string(),
231 requirements: vec![],
232 }
233 }
234
235 fn exposure(&self) -> ToolExposure {
236 TypedTool::exposure(self)
237 }
238
239 fn should_activate(&self, ctx: &ActivationContext) -> bool {
240 TypedTool::should_activate(self, ctx)
241 }
242
243 async fn call(&self, args: &Value, ctx: &ToolContext) -> AgentResult<Vec<Content>> {
244 let typed_args: T::Args =
245 serde_json::from_value(args.clone()).map_err(|_| AgentError::ToolArgsInvalid {
246 name: self.name().to_string(),
247 raw: args.to_string(),
248 })?;
249 let output = self.call_typed(typed_args, ctx).await?;
250 Ok(vec![self.format_output(output)])
251 }
252}
253
254pub fn render_tool_definition(tool: &dyn Tool) -> Value {
259 json!({
260 "type": "function",
261 "function": {
262 "name": tool.name(),
263 "description": tool.description(),
264 "parameters": tool.schema(),
265 }
266 })
267}
268
269pub(crate) type ToolRef = Arc<dyn Tool>;
270
271#[derive(Clone, Default)]
272pub struct ToolRegistry {
273 tools: HashMap<String, ToolRef>,
274}
275
276impl ToolRegistry {
277 pub fn register(&mut self, tool: impl Tool + 'static) {
278 self.tools.insert(tool.name().to_string(), Arc::new(tool));
279 }
280
281 pub fn register_arc(&mut self, tool: Arc<dyn Tool>) {
282 self.tools.insert(tool.name().to_string(), tool);
283 }
284
285 pub fn remove(&mut self, name: &str) {
287 self.tools.remove(name);
288 }
289
290 pub fn get(&self, name: &str) -> Option<ToolRef> {
291 self.tools.get(name).cloned()
292 }
293
294 pub fn definitions(&self) -> Vec<Value> {
295 let mut tools: Vec<_> = self.tools.values().collect();
296 tools.sort_by_key(|t| t.name());
297 tools
298 .into_iter()
299 .map(|t| render_tool_definition(t.as_ref()))
300 .collect()
301 }
302
303 pub fn definitions_filtered(&self, ctx: &ActivationContext) -> Vec<Value> {
309 let mut tools: Vec<_> = self.tools.values().collect();
310 tools.sort_by_key(|t| t.name());
311
312 let direct_names: Vec<String> = tools
318 .iter()
319 .filter(|t| t.exposure() == ToolExposure::Direct)
320 .map(|t| t.name().to_string())
321 .collect();
322
323 let mut activated_names = direct_names.clone();
324 for t in &tools {
325 if t.exposure() == ToolExposure::Deferred {
326 let mut ctx_with_tools = ctx.clone();
327 ctx_with_tools.current_tools = activated_names.clone();
328 if t.should_activate(&ctx_with_tools) {
329 activated_names.push(t.name().to_string());
330 }
331 }
332 }
333
334 tools
335 .into_iter()
336 .filter(|t| match t.exposure() {
337 ToolExposure::Direct => true,
338 ToolExposure::Deferred => activated_names.contains(&t.name().to_string()),
339 ToolExposure::Hidden => false,
340 })
341 .map(|t| render_tool_definition(t.as_ref()))
342 .collect()
343 }
344
345 pub fn len(&self) -> usize {
346 self.tools.len()
347 }
348
349 pub fn is_empty(&self) -> bool {
350 self.tools.is_empty()
351 }
352
353 pub fn metadatas(&self) -> Vec<ToolMetadata> {
359 let mut list: Vec<_> = self.tools.values().map(|tool| tool.metadata()).collect();
360 list.sort_by(|a, b| a.name.cmp(&b.name));
361 list
362 }
363}
364
365#[cfg(test)]
366mod tests {
367 use super::*;
368
369 #[test]
370 fn content_text_ctor_and_into_vec() {
371 let c = Content::text("hello");
372 let v: Vec<Content> = c.clone().into();
373 assert_eq!(v.len(), 1);
374 assert!(matches!(v[0], Content::Text { .. }));
375 assert!(matches!(&c, Content::Text { text } if text == "hello"));
376 }
377
378 #[test]
379 fn content_serializes_with_type_tag() {
380 let c = Content::text("hi");
381 let j = serde_json::to_value(&c).unwrap();
382 assert_eq!(j["type"], "text");
383 assert_eq!(j["text"], "hi");
384 }
385
386 #[test]
387 fn tool_context_for_test_constructs() {
388 let ctx = ToolContext::for_test();
389 assert!(ctx.llm_client.is_none());
390 assert!(ctx.session_store.is_none());
391 assert!(!ctx.cancel_token.is_cancelled());
392 ctx.emit_progress("hello");
393 }
394
395 #[test]
396 fn typed_tool_schema_is_derived_from_args() {
397 #[derive(schemars::JsonSchema, serde::Deserialize, serde::Serialize)]
398 struct GreetArgs {
399 name: String,
400 #[serde(default)]
401 times: u32,
402 }
403
404 let schema = schemars::schema_for!(GreetArgs);
405 let j = serde_json::to_value(&schema).unwrap();
406 assert!(j["properties"]["name"].is_object());
408 assert!(j["properties"]["times"].is_object());
409 }
410
411 #[test]
412 fn typed_tool_schema_is_provider_safe_for_nested_enum() {
413 #[derive(schemars::JsonSchema, serde::Deserialize, serde::Serialize)]
414 enum Status {
415 Active,
416 Paused,
417 }
418
419 #[derive(schemars::JsonSchema, serde::Deserialize, serde::Serialize)]
420 struct Args {
421 name: String,
422 status: Status,
423 }
424
425 #[derive(Default)]
426 struct NestedTool;
427 #[async_trait]
428 impl TypedTool for NestedTool {
429 type Args = Args;
430 type Output = String;
431 fn name(&self) -> &'static str {
432 "nested"
433 }
434 fn description(&self) -> &'static str {
435 ""
436 }
437 async fn call_typed(
438 &self,
439 _args: Args,
440 _ctx: &ToolContext,
441 ) -> crate::types::AgentResult<String> {
442 Ok(String::new())
443 }
444 }
445
446 let schema = Tool::schema(&NestedTool);
447 let raw = schema.to_string();
448 assert!(!raw.contains("$ref"), "schema contains $ref: {raw}");
451 assert!(!raw.contains("$defs"), "schema contains $defs: {raw}");
452 assert!(
453 !raw.contains("definitions"),
454 "schema has definitions: {raw}"
455 );
456 assert!(schema.get("$schema").is_none(), "schema has $schema key");
457
458 let variants: Vec<&str> = schema["properties"]["status"]["enum"]
460 .as_array()
461 .unwrap()
462 .iter()
463 .map(|v| v.as_str().unwrap())
464 .collect();
465 assert!(variants.contains(&"Active"), "missing Active: {variants:?}");
466 assert!(variants.contains(&"Paused"), "missing Paused: {variants:?}");
467 }
468
469 #[test]
470 fn definitions_are_sorted_by_name() {
471 struct NamedTool(&'static str);
472 #[async_trait::async_trait]
473 impl Tool for NamedTool {
474 fn name(&self) -> &'static str {
475 self.0
476 }
477 fn description(&self) -> &'static str {
478 ""
479 }
480 fn schema(&self) -> serde_json::Value {
481 serde_json::Value::Null
482 }
483 async fn call(
484 &self,
485 _args: &serde_json::Value,
486 _ctx: &ToolContext,
487 ) -> crate::types::AgentResult<Vec<Content>> {
488 Ok(vec![])
489 }
490 }
491
492 let mut registry = ToolRegistry::default();
493 registry.register(NamedTool("zeta"));
494 registry.register(NamedTool("alpha"));
495 registry.register(NamedTool("mike"));
496
497 let defs = registry.definitions();
498 let names: Vec<&str> = defs
499 .iter()
500 .map(|d| d["function"]["name"].as_str().unwrap())
501 .collect();
502 assert_eq!(names, vec!["alpha", "mike", "zeta"]);
503 }
504
505 #[test]
508 fn content_image_and_content_text_skips_images() {
509 let img = Content::image("base64data", "image/png");
510 assert!(
511 matches!(&img, Content::Image { data, mime_type } if data == "base64data" && mime_type == "image/png")
512 );
513
514 let text = content_text(&[
515 Content::text("a"),
516 Content::image("b", "image/png"),
517 Content::text("c"),
518 ]);
519 assert_eq!(text, "a\nc");
520 }
521
522 #[test]
523 fn emit_partial_result_sends_event() {
524 let (tx, mut rx) = mpsc::unbounded_channel();
525 let ctx = ToolContext {
526 session_id: SessionId::new(0),
527 user_event_tx: tx,
528 llm_client: None,
529 session_store: None,
530 language: crate::types::Language::En,
531 cancel_token: tokio_util::sync::CancellationToken::new(),
532 max_output_chars: None,
533 event_bus: crate::engine::EventBus::new(1),
534 };
535 ctx.emit_partial_result("tc1", "partial", true);
536 match rx.try_recv().unwrap() {
537 UserEvent::ToolPartialResult {
538 tool_call_id,
539 content,
540 is_partial,
541 } => {
542 assert_eq!(tool_call_id, "tc1");
543 assert_eq!(content, "partial");
544 assert!(is_partial);
545 }
546 other => panic!("unexpected event: {other:?}"),
547 }
548 }
549
550 #[test]
556 fn emit_user_event_dual_writes_to_event_bus() {
557 let (tx, mut rx) = mpsc::unbounded_channel();
558 let ctx = ToolContext {
559 session_id: SessionId::new(7),
560 user_event_tx: tx,
561 llm_client: None,
562 session_store: None,
563 language: crate::types::Language::En,
564 cancel_token: tokio_util::sync::CancellationToken::new(),
565 max_output_chars: None,
566 event_bus: crate::engine::EventBus::new(4),
567 };
568 let mut bus_rx = ctx.event_bus.subscribe();
569 ctx.emit_progress("working");
570
571 assert!(
573 matches!(rx.try_recv().unwrap(), UserEvent::Progress { text } if text == "working")
574 );
575
576 match bus_rx.try_recv().unwrap() {
578 crate::types::RuntimeEvent::UserEvent {
579 session_id,
580 event,
581 agent_id,
582 ..
583 } => {
584 assert_eq!(session_id, SessionId::new(7));
585 assert!(matches!(event, UserEvent::Progress { text } if text == "working"));
586 assert!(agent_id.is_none());
587 }
588 other => panic!("unexpected bus event: {other:?}"),
589 }
590 }
591
592 #[derive(schemars::JsonSchema, serde::Deserialize, serde::Serialize)]
594 struct GreetArgs {
595 name: String,
596 }
597
598 struct GreetTool;
599 #[async_trait]
600 impl TypedTool for GreetTool {
601 type Args = GreetArgs;
602 type Output = String;
603 fn name(&self) -> &'static str {
604 "greet"
605 }
606 fn description(&self) -> &'static str {
607 "greets a name"
608 }
609 fn origin(&self) -> &'static str {
610 "test-crate"
611 }
612 fn version(&self) -> &'static str {
613 "1.0.0"
614 }
615 async fn call_typed(&self, args: GreetArgs, _ctx: &ToolContext) -> AgentResult<String> {
616 Ok(format!("Hello, {}!", args.name))
617 }
618 }
619
620 #[test]
621 fn typed_tool_blanket_delegates_name_description() {
622 let t = GreetTool;
623 assert_eq!(Tool::name(&t), "greet");
624 assert_eq!(Tool::description(&t), "greets a name");
625 }
626
627 #[test]
628 fn typed_tool_metadata_uses_origin_and_version() {
629 let m = Tool::metadata(&GreetTool);
630 assert_eq!(m.name, "greet");
631 assert_eq!(m.description, "greets a name");
632 assert_eq!(m.origin, "test-crate");
633 assert_eq!(m.version, "1.0.0");
634 assert!(m.requirements.is_empty());
635 }
636
637 #[tokio::test]
638 async fn typed_tool_call_deserializes_and_formats() {
639 let ctx = ToolContext::for_test();
640 let out = Tool::call(&GreetTool, &json!({"name": "world"}), &ctx)
641 .await
642 .unwrap();
643 assert_eq!(content_text(&out), "Hello, world!");
645 }
646
647 #[derive(serde::Serialize)]
649 struct GreetResult {
650 message: String,
651 }
652
653 struct GreetStructTool;
654 #[async_trait]
655 impl TypedTool for GreetStructTool {
656 type Args = GreetArgs;
657 type Output = GreetResult;
658 fn name(&self) -> &'static str {
659 "greet_struct"
660 }
661 fn description(&self) -> &'static str {
662 "greets as json"
663 }
664 async fn call_typed(
665 &self,
666 args: GreetArgs,
667 _ctx: &ToolContext,
668 ) -> AgentResult<GreetResult> {
669 Ok(GreetResult {
670 message: format!("Hello, {}!", args.name),
671 })
672 }
673 }
674
675 #[test]
676 fn format_output_json_serializes_struct() {
677 let out = GreetStructTool.format_output(GreetResult {
678 message: "hi".into(),
679 });
680 assert_eq!(content_text(&[out]), r#"{"message":"hi"}"#);
681 }
682
683 #[tokio::test]
684 async fn typed_tool_call_invalid_args_is_tool_args_invalid() {
685 let ctx = ToolContext::for_test();
686 let err = Tool::call(&GreetTool, &json!({"nope": 1}), &ctx)
687 .await
688 .unwrap_err();
689 assert!(matches!(err, AgentError::ToolArgsInvalid { .. }));
690 }
691
692 struct NamedTool(&'static str);
693 #[async_trait]
694 impl Tool for NamedTool {
695 fn name(&self) -> &'static str {
696 self.0
697 }
698 fn description(&self) -> &'static str {
699 ""
700 }
701 fn schema(&self) -> serde_json::Value {
702 serde_json::Value::Null
703 }
704 async fn call(
705 &self,
706 _args: &serde_json::Value,
707 _ctx: &ToolContext,
708 ) -> AgentResult<Vec<Content>> {
709 Ok(vec![])
710 }
711 }
712
713 #[test]
714 fn registry_register_arc_get_remove_len_is_empty() {
715 let mut r = ToolRegistry::default();
716 assert!(r.is_empty());
717 assert_eq!(r.len(), 0);
718
719 let t: Arc<dyn Tool> = Arc::new(NamedTool("x"));
720 r.register_arc(t);
721 assert!(!r.is_empty());
722 assert_eq!(r.len(), 1);
723 assert!(r.get("x").is_some());
724 assert!(r.get("missing").is_none());
725
726 r.remove("x");
727 assert!(r.is_empty());
728 }
729
730 #[test]
731 fn metadatas_are_sorted_by_name() {
732 let mut r = ToolRegistry::default();
733 r.register(NamedTool("zeta"));
734 r.register(NamedTool("alpha"));
735
736 let metas = r.metadatas();
737 let names: Vec<&str> = metas.iter().map(|m| m.name.as_str()).collect();
738 assert_eq!(names, vec!["alpha", "zeta"]);
739 assert_eq!(metas[0].origin, "custom");
740 assert_eq!(metas[0].version, "unknown");
741 }
742
743 struct ExposureTool(&'static str, ToolExposure);
746 #[async_trait]
747 impl Tool for ExposureTool {
748 fn name(&self) -> &'static str {
749 self.0
750 }
751 fn description(&self) -> &'static str {
752 ""
753 }
754 fn schema(&self) -> serde_json::Value {
755 serde_json::Value::Null
756 }
757 async fn call(
758 &self,
759 _args: &serde_json::Value,
760 _ctx: &ToolContext,
761 ) -> AgentResult<Vec<Content>> {
762 Ok(vec![])
763 }
764 fn exposure(&self) -> ToolExposure {
765 self.1.clone()
766 }
767 }
768
769 struct ConditionalTool(&'static str, bool);
770 #[async_trait]
771 impl Tool for ConditionalTool {
772 fn name(&self) -> &'static str {
773 self.0
774 }
775 fn description(&self) -> &'static str {
776 ""
777 }
778 fn schema(&self) -> serde_json::Value {
779 serde_json::Value::Null
780 }
781 async fn call(
782 &self,
783 _args: &serde_json::Value,
784 _ctx: &ToolContext,
785 ) -> AgentResult<Vec<Content>> {
786 Ok(vec![])
787 }
788 fn exposure(&self) -> ToolExposure {
789 ToolExposure::Deferred
790 }
791 fn should_activate(&self, _ctx: &ActivationContext) -> bool {
792 self.1
793 }
794 }
795
796 fn default_ctx() -> ActivationContext {
797 ActivationContext {
798 session_id: crate::types::SessionId::new(0),
799 current_tools: vec![],
800 workspace: std::path::PathBuf::from("/tmp"),
801 }
802 }
803
804 #[test]
805 fn tool_exposure_default_is_direct() {
806 let t = NamedTool("x");
808 assert_eq!(t.exposure(), ToolExposure::Direct);
809 }
810
811 #[test]
812 fn tool_should_activate_default_is_true() {
813 let t = NamedTool("x");
814 assert!(t.should_activate(&default_ctx()));
815 }
816
817 #[test]
818 fn definitions_filtered_includes_direct_excludes_hidden() {
819 let mut r = ToolRegistry::default();
820 r.register(ExposureTool("direct_a", ToolExposure::Direct));
821 r.register(ExposureTool("hidden_a", ToolExposure::Hidden));
822 r.register(ExposureTool("direct_b", ToolExposure::Direct));
823
824 let defs = r.definitions_filtered(&default_ctx());
825 let names: Vec<&str> = defs
826 .iter()
827 .map(|d| d["function"]["name"].as_str().unwrap())
828 .collect();
829 assert_eq!(names, vec!["direct_a", "direct_b"]);
830 }
831
832 #[test]
833 fn definitions_filtered_includes_deferred_when_activated() {
834 let mut r = ToolRegistry::default();
835 r.register(ExposureTool("always", ToolExposure::Direct));
836 r.register(ConditionalTool("maybe", true)); r.register(ConditionalTool("never", false)); let defs = r.definitions_filtered(&default_ctx());
840 let names: Vec<&str> = defs
841 .iter()
842 .map(|d| d["function"]["name"].as_str().unwrap())
843 .collect();
844 assert_eq!(names, vec!["always", "maybe"]);
845 }
846
847 #[test]
848 fn definitions_filtered_all_hidden_returns_empty() {
849 let mut r = ToolRegistry::default();
850 r.register(ExposureTool("h1", ToolExposure::Hidden));
851 r.register(ExposureTool("h2", ToolExposure::Hidden));
852
853 let defs = r.definitions_filtered(&default_ctx());
854 assert!(defs.is_empty());
855 }
856
857 #[test]
858 fn definitions_filtered_empty_registry() {
859 let r = ToolRegistry::default();
860 let defs = r.definitions_filtered(&default_ctx());
861 assert!(defs.is_empty());
862 }
863
864 #[test]
865 fn definitions_unfiltered_includes_hidden() {
866 let mut r = ToolRegistry::default();
868 r.register(ExposureTool("visible", ToolExposure::Direct));
869 r.register(ExposureTool("secret", ToolExposure::Hidden));
870
871 let defs = r.definitions();
872 assert_eq!(defs.len(), 2);
873 }
874}