1use 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
21pub 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 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 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 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 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 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 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 #[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 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
295fn 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}