Skip to main content

ares_tools/
tool_service.rs

1//! Tools capability: tenant-aware tool resolution.
2//!
3//! Precedence: `tenant runtime → fleet runtime → static` (static includes
4//! `mcp_bridge` registrations). Callers obtain the service via
5//! `ctx.get::<Tools>()` and isolate with `ctx.isolate::<Tools>(tenant_id)`.
6
7use std::any::TypeId;
8use std::collections::HashSet;
9use std::future::Future;
10use std::sync::Arc;
11
12use ares_types::types::{Result, ToolDefinition};
13use cordis::{CordisError, EventsService, Service};
14use serde_json::{json, Value};
15
16use crate::registry::{Tool, ToolRegistry};
17
18#[cfg(any(feature = "postgres", test))]
19use crate::runtime_registry::RuntimeToolRegistry;
20
21/// Tenant-aware tool capability.
22///
23/// Isolate labels on [`Tools`] win over a `TenantContext` intercept.
24pub struct Tools {
25    static_registry: Arc<ToolRegistry>,
26    #[cfg(any(feature = "postgres", test))]
27    runtime: Option<Arc<RuntimeToolRegistry>>,
28}
29
30impl Clone for Tools {
31    fn clone(&self) -> Self {
32        Self {
33            static_registry: Arc::clone(&self.static_registry),
34            #[cfg(any(feature = "postgres", test))]
35            runtime: self.runtime.clone(),
36        }
37    }
38}
39
40impl Tools {
41    pub(crate) fn new(static_registry: Arc<ToolRegistry>) -> Self {
42        Self {
43            static_registry,
44            #[cfg(any(feature = "postgres", test))]
45            runtime: None,
46        }
47    }
48
49    #[cfg(any(feature = "postgres", test))]
50    #[allow(dead_code)]
51    pub(crate) fn with_runtime(
52        static_registry: Arc<ToolRegistry>,
53        runtime: Option<Arc<RuntimeToolRegistry>>,
54    ) -> Self {
55        Self {
56            static_registry,
57            runtime,
58        }
59    }
60
61    /// Build Tools from a static tool set. Runtime is unset.
62    pub fn from_static(tools: impl IntoIterator<Item = Arc<dyn Tool>>) -> Self {
63        let mut registry = ToolRegistry::new();
64        for tool in tools {
65            registry.register(tool);
66        }
67        Self::new(Arc::new(registry))
68    }
69
70    /// Validate runtime tool execution_config without exposing the registry.
71    pub fn validate_runtime_tool_execution_config(
72        tool_type: &str,
73        execution_config: &Value,
74    ) -> Result<()> {
75        #[cfg(any(feature = "postgres", test))]
76        {
77            RuntimeToolRegistry::validate_execution_config(tool_type, execution_config)
78        }
79        #[cfg(not(any(feature = "postgres", test)))]
80        {
81            let _ = (tool_type, execution_config);
82            Ok(())
83        }
84    }
85
86    /// Resolve a tool using the tenant derived from `ctx` (isolate, then intercept).
87    pub fn resolve(&self, ctx: &Arc<cordis::Context>, name: &str) -> Option<Arc<dyn Tool>> {
88        let tenant = tenant_id_from_tool_ctx(ctx);
89        let Some(events) = ctx.get::<EventsService>() else {
90            return self.resolve_named(name, tenant.as_deref());
91        };
92        let payload = serde_json::to_value(cordis::ToolsResolveRequest {
93            name: name.to_string(),
94            tenant: tenant.clone(),
95        })
96        .unwrap_or(serde_json::Value::Null);
97        let this = self.clone();
98        let out = match run_waterfall(&events, "tools.resolve", payload, move |p| async move {
99            let n = p
100                .get("name")
101                .and_then(Value::as_str)
102                .unwrap_or("")
103                .to_string();
104            let tenant = p.get("tenant").and_then(Value::as_str).map(str::to_string);
105            let found = this.resolve_named(&n, tenant.as_deref()).is_some();
106            Ok(json!({
107                "name": n,
108                "tenant": p.get("tenant").cloned().unwrap_or(Value::Null),
109                "found": found,
110            }))
111        }) {
112            Ok(v) => v,
113            Err(_) => return self.resolve_named(name, tenant.as_deref()),
114        };
115        if out.get("deny").and_then(Value::as_bool) == Some(true) {
116            return None;
117        }
118        out.get("found")?;
119        let resolved_name = out.get("name").and_then(Value::as_str).unwrap_or(name);
120        self.resolve_named(resolved_name, tenant.as_deref())
121    }
122
123    /// List tools using the tenant derived from `ctx` (isolate, then intercept).
124    pub fn list(&self, ctx: &Arc<cordis::Context>) -> Vec<ToolDefinition> {
125        let tenant = tenant_id_from_tool_ctx(ctx);
126        let Some(events) = ctx.get::<EventsService>() else {
127            return self.list_named(tenant.as_deref());
128        };
129        let payload = serde_json::to_value(cordis::ToolsListRequest {
130            tenant: tenant.clone(),
131        })
132        .unwrap_or(serde_json::Value::Null);
133        let this = self.clone();
134        let out = match run_waterfall(&events, "tools.list", payload, move |p| async move {
135            let tenant = p.get("tenant").and_then(Value::as_str).map(str::to_string);
136            let tools = this.list_named(tenant.as_deref());
137            Ok(json!({
138                "tenant": p.get("tenant").cloned().unwrap_or(Value::Null),
139                "tools": tools,
140            }))
141        }) {
142            Ok(v) => v,
143            Err(_) => return self.list_named(tenant.as_deref()),
144        };
145        match out
146            .get("tools")
147            .cloned()
148            .and_then(|t| serde_json::from_value::<Vec<ToolDefinition>>(t).ok())
149        {
150            Some(defs) => defs,
151            None => self.list_named(tenant.as_deref()),
152        }
153    }
154
155    /// Execute a named tool, wrapping the call in `tools.execute` around-middleware
156    /// when [`EventsService`] is on `ctx`.
157    pub async fn execute(
158        &self,
159        ctx: &Arc<cordis::Context>,
160        name: &str,
161        args: Value,
162    ) -> Result<Value> {
163        let tool = self
164            .resolve(ctx, name)
165            .ok_or_else(|| ares_types::AppError::NotFound(format!("Tool not found: {name}")))?;
166        let Some(events) = ctx.get::<EventsService>() else {
167            return tool.execute(args).await;
168        };
169        let payload = serde_json::to_value(cordis::ToolsExecutePayload {
170            name: name.to_string(),
171            args,
172        })
173        .unwrap_or(serde_json::Value::Null);
174        let out = events
175            .waterfall_around(
176                cordis::events_catalog::ev::TOOLS_EXECUTE.to_string(),
177                payload,
178                move |p| async move {
179                    let exec_args = p.get("args").cloned().unwrap_or(Value::Null);
180                    let result = tool
181                        .execute(exec_args)
182                        .await
183                        .map_err(|e| CordisError::Fiber(e.to_string()))?;
184                    let mut out = p;
185                    if let Some(obj) = out.as_object_mut() {
186                        obj.insert("result".into(), result);
187                    } else {
188                        out = json!({ "result": result });
189                    }
190                    Ok(out)
191                },
192            )
193            .await
194            .map_err(|e| ares_types::AppError::Internal(e.to_string()))?;
195        Ok(out.get("result").cloned().unwrap_or(Value::Null))
196    }
197
198    /// Reload runtime tools from the database when a runtime registry is attached.
199    pub async fn reload(&self) -> Result<()> {
200        #[cfg(any(feature = "postgres", test))]
201        if let Some(rt) = &self.runtime {
202            rt.reload().await?;
203        }
204        Ok(())
205    }
206
207    /// Runtime registry for admin mutation. Not a Service; do not `ctx.get` it.
208    #[cfg(any(feature = "postgres", test))]
209    #[allow(dead_code)]
210    pub(crate) fn runtime(&self) -> Option<Arc<RuntimeToolRegistry>> {
211        self.runtime.clone()
212    }
213
214    /// Concrete runtime tool type after tenant visibility checks.
215    pub fn tool_type(&self, ctx: &Arc<cordis::Context>, name: &str) -> Option<String> {
216        #[cfg(any(feature = "postgres", test))]
217        {
218            let tenant = tenant_id_from_tool_ctx(ctx);
219            self.runtime
220                .as_ref()
221                .and_then(|rt| rt.tool_type_for_tenant(name, tenant.as_deref()))
222        }
223        #[cfg(not(any(feature = "postgres", test)))]
224        {
225            let _ = (ctx, name);
226            None
227        }
228    }
229
230    fn resolve_named(&self, name: &str, tenant: Option<&str>) -> Option<Arc<dyn Tool>> {
231        #[cfg(any(feature = "postgres", test))]
232        if let Some(rt) = &self.runtime {
233            if let Some(tid) = tenant {
234                if let Some(tool) = rt.get_for_tenant(name, Some(tid)) {
235                    return Some(tool);
236                }
237            }
238            if let Some(tool) = rt.get(name) {
239                return Some(tool);
240            }
241        }
242        #[cfg(not(any(feature = "postgres", test)))]
243        let _ = tenant;
244        self.static_registry.get(name).cloned()
245    }
246
247    fn list_named(&self, tenant: Option<&str>) -> Vec<ToolDefinition> {
248        let mut seen: HashSet<String> = HashSet::new();
249        let mut out: Vec<ToolDefinition> = Vec::new();
250
251        let push_defs = |defs: Vec<ToolDefinition>,
252                         seen: &mut HashSet<String>,
253                         out: &mut Vec<ToolDefinition>| {
254            for d in defs {
255                if seen.insert(d.name.clone()) {
256                    out.push(d);
257                }
258            }
259        };
260
261        #[cfg(any(feature = "postgres", test))]
262        if let Some(rt) = &self.runtime {
263            push_defs(
264                rt.get_tool_definitions_for_tenant(tenant),
265                &mut seen,
266                &mut out,
267            );
268            if tenant.is_some() {
269                let remaining: Vec<ToolDefinition> = rt
270                    .get_tool_definitions()
271                    .into_iter()
272                    .filter(|d| !seen.contains(&d.name))
273                    .collect();
274                push_defs(remaining, &mut seen, &mut out);
275            }
276        }
277        #[cfg(not(any(feature = "postgres", test)))]
278        let _ = tenant;
279
280        push_defs(
281            self.static_registry.get_tool_definitions(),
282            &mut seen,
283            &mut out,
284        );
285        out
286    }
287}
288
289impl Service for Tools {
290    fn check(&self) -> bool {
291        true
292    }
293}
294
295/// Derive the tenant id for tool resolution from `ctx`.
296///
297/// Isolate labels on [`Tools`] win. A leading `tenant:` or `user:` prefix is
298/// stripped; a non-empty remainder is the tenant. If the isolate label is
299/// missing or empty after stripping, fall back to a `TenantContext` intercept.
300/// Unlabeled contexts with no intercept yield `None`.
301fn tenant_id_from_tool_ctx(ctx: &Arc<cordis::Context>) -> Option<String> {
302    if let Some(label) = ctx.isolate_label(TypeId::of::<Tools>()) {
303        let trimmed = label
304            .strip_prefix("tenant:")
305            .or_else(|| label.strip_prefix("user:"))
306            .unwrap_or(&label);
307        if !trimmed.is_empty() {
308            return Some(trimmed.to_string());
309        }
310    }
311    ctx.get::<ares_types::models::TenantContext>()
312        .map(|tc| tc.tenant_id.clone())
313        .filter(|id| !id.is_empty())
314}
315
316fn run_waterfall<F, Fut>(
317    events: &EventsService,
318    event: &str,
319    payload: Value,
320    core: F,
321) -> std::result::Result<Value, CordisError>
322where
323    F: FnOnce(Value) -> Fut + Send + 'static,
324    Fut: Future<Output = std::result::Result<Value, CordisError>> + Send + 'static,
325{
326    let Ok(handle) = tokio::runtime::Handle::try_current() else {
327        return Err(CordisError::Fiber("no tokio runtime".into()));
328    };
329    tokio::task::block_in_place(|| {
330        handle.block_on(events.waterfall_around(event.into(), payload, core))
331    })
332}
333
334#[cfg(test)]
335mod tests {
336    use super::*;
337    use ares_types::models::{TenantContext, TenantTier};
338    use async_trait::async_trait;
339    use cordis::Context;
340    use std::sync::atomic::{AtomicBool, Ordering};
341
342    struct ProbeTool {
343        name: String,
344        ran: Option<Arc<AtomicBool>>,
345    }
346
347    impl ProbeTool {
348        fn new(name: impl Into<String>) -> Self {
349            Self {
350                name: name.into(),
351                ran: None,
352            }
353        }
354    }
355
356    #[async_trait]
357    impl Tool for ProbeTool {
358        fn name(&self) -> &str {
359            &self.name
360        }
361        fn description(&self) -> &str {
362            "probe"
363        }
364        fn parameters_schema(&self) -> Value {
365            json!({})
366        }
367        async fn execute(&self, _args: Value) -> Result<Value> {
368            if let Some(ran) = &self.ran {
369                ran.store(true, Ordering::SeqCst);
370            }
371            Ok(json!({ "ok": self.name }))
372        }
373    }
374
375    #[test]
376    fn unlabeled_root_yields_no_tenant() {
377        let ctx = Context::new_root();
378        assert_eq!(tenant_id_from_tool_ctx(&ctx), None);
379    }
380
381    #[test]
382    fn intercept_tenant_context_yields_acme() {
383        let ctx =
384            Context::new_root().with_intercept(TenantContext::new("acme".into(), TenantTier::Pro));
385        assert_eq!(tenant_id_from_tool_ctx(&ctx).as_deref(), Some("acme"));
386    }
387
388    #[test]
389    fn isolate_wins_over_intercept() {
390        let ctx = Context::new_root()
391            .with_intercept(TenantContext::new("acme".into(), TenantTier::Pro))
392            .isolate::<Tools>("tenant:iso");
393        assert_eq!(tenant_id_from_tool_ctx(&ctx).as_deref(), Some("iso"));
394    }
395
396    #[test]
397    fn resolve_missing_tool_is_none() {
398        let svc = Tools::new(Arc::new(ToolRegistry::new()));
399        let ctx = Context::new_root();
400        assert!(svc.resolve(&ctx, "missing").is_none());
401        assert!(svc.list(&ctx).is_empty());
402    }
403
404    #[test]
405    fn list_and_resolve_use_ctx_isolate() {
406        let mut registry = ToolRegistry::new();
407        registry.register(Arc::new(crate::calculator::Calculator));
408        let svc = Tools::with_runtime(Arc::new(registry), None);
409        let ctx = Context::new_root().isolate::<Tools>("tenant:acme");
410        assert!(svc.resolve(&ctx, "calculator").is_some());
411        assert!(svc.list(&ctx).iter().any(|d| d.name == "calculator"));
412        assert!(svc.resolve(&ctx, "unknown").is_none());
413    }
414
415    #[test]
416    fn from_static_resolves_calculator() {
417        let svc = Tools::from_static([Arc::new(crate::calculator::Calculator) as Arc<dyn Tool>]);
418        let ctx = Context::new_root();
419        assert!(svc.resolve(&ctx, "calculator").is_some());
420    }
421
422    #[tokio::test(flavor = "multi_thread")]
423    async fn tools_list_waterfall_filters_tool() {
424        let svc = Tools::from_static([
425            Arc::new(ProbeTool::new("a")) as Arc<dyn Tool>,
426            Arc::new(ProbeTool::new("b")) as Arc<dyn Tool>,
427        ]);
428        let ctx = Context::new_root();
429        ctx.provide(EventsService::new());
430        let names: Vec<_> = svc.list(&ctx).into_iter().map(|d| d.name).collect();
431        assert!(names.contains(&"a".to_string()));
432        assert!(names.contains(&"b".to_string()));
433
434        let events = ctx.get::<EventsService>().expect("events");
435        events.on_waterfall(
436            cordis::events_catalog::ev::TOOLS_LIST.to_string(),
437            |payload, next| async move {
438                let mut out = next(payload).await?;
439                if let Some(arr) = out.get_mut("tools").and_then(Value::as_array_mut) {
440                    arr.retain(|t| t.get("name").and_then(Value::as_str) != Some("b"));
441                }
442                Ok(out)
443            },
444        );
445        let names: Vec<_> = svc.list(&ctx).into_iter().map(|d| d.name).collect();
446        assert!(names.contains(&"a".to_string()));
447        assert!(!names.contains(&"b".to_string()));
448    }
449
450    #[tokio::test(flavor = "multi_thread")]
451    async fn tools_resolve_waterfall_deny() {
452        let svc = Tools::from_static([Arc::new(ProbeTool::new("a")) as Arc<dyn Tool>]);
453        let ctx = Context::new_root();
454        ctx.provide(EventsService::new());
455        assert!(svc.resolve(&ctx, "a").is_some());
456        let events = ctx.get::<EventsService>().expect("events");
457        events.on_waterfall(
458            cordis::events_catalog::ev::TOOLS_RESOLVE.to_string(),
459            |payload, _next| async move { Ok(payload) },
460        );
461        assert!(svc.resolve(&ctx, "a").is_none());
462    }
463
464    #[tokio::test(flavor = "multi_thread")]
465    async fn tools_execute_core_runs() {
466        let svc = Tools::from_static([Arc::new(ProbeTool::new("probe")) as Arc<dyn Tool>]);
467        let ctx = Context::new_root();
468        ctx.provide(EventsService::new());
469        let out = svc
470            .execute(&ctx, "probe", json!({}))
471            .await
472            .expect("execute");
473        assert_eq!(out, json!({ "ok": "probe" }));
474    }
475
476    #[tokio::test(flavor = "multi_thread")]
477    async fn tools_execute_short_circuit_skips_tool() {
478        let ran = Arc::new(AtomicBool::new(false));
479        let svc = Tools::from_static([Arc::new(ProbeTool {
480            name: "probe".into(),
481            ran: Some(Arc::clone(&ran)),
482        }) as Arc<dyn Tool>]);
483        let ctx = Context::new_root();
484        ctx.provide(EventsService::new());
485        let events = ctx.get::<EventsService>().expect("events");
486        events.on_waterfall(
487            cordis::events_catalog::ev::TOOLS_EXECUTE.to_string(),
488            |_payload, _next| async move { Ok(json!({ "result": { "short": true } })) },
489        );
490        let out = svc
491            .execute(&ctx, "probe", json!({}))
492            .await
493            .expect("execute");
494        assert_eq!(out, json!({ "short": true }));
495        assert!(
496            !ran.load(Ordering::SeqCst),
497            "tool.execute must not run when handler short-circuits"
498        );
499    }
500}