Skip to main content

ares_tools/
plugins.rs

1//! Loader plugin factories for CalculatorService and Tools.
2//!
3//! Factories call `cordis::Context::plugin` via block_in_place + block_on.
4//! ToolRegistry and RuntimeToolRegistry are not provided as Services.
5
6use std::collections::HashMap;
7use std::sync::Arc;
8
9use serde_json::Value;
10
11use crate::config::ToolConfig;
12use crate::registry::ToolRegistry;
13use crate::{Calculator, CalculatorConfig, CalculatorService, Tools};
14
15fn block_on_plugin<S: cordis::Service + 'static>(
16    ctx: &Arc<cordis::Context>,
17    svc: S,
18) -> Result<cordis::FiberId, cordis::CordisError> {
19    tokio::task::block_in_place(|| tokio::runtime::Handle::current().block_on(ctx.plugin(svc)))
20}
21
22/// Register the `CalculatorService` and `Tools` loader factories
23/// (manual fallback path; inventory carries the same pair).
24pub fn register_plugins(reg: &cordis::PluginRegistry) {
25    reg.register("CalculatorService", Arc::new(factory_calculator));
26    reg.register("Tools", Arc::new(factory_tools));
27}
28
29#[cfg(feature = "inventory")]
30inventory::submit! {
31    cordis::CordisPluginFactory { name: "CalculatorService", make: factory_calculator }
32}
33
34#[cfg(feature = "inventory")]
35inventory::submit! {
36    cordis::CordisPluginFactory { name: "Tools", make: factory_tools }
37}
38
39fn factory_calculator(
40    ctx: &Arc<cordis::Context>,
41    config: &Value,
42) -> Result<cordis::FiberId, cordis::CordisError> {
43    let calculator_config = if config.is_null()
44        || config.as_object().is_some_and(|object| object.is_empty())
45    {
46        CalculatorConfig
47    } else {
48        serde_json::from_value::<CalculatorConfig>(config.clone()).map_err(|error| {
49            cordis::CordisError::Configuration(format!("invalid CalculatorService config: {error}"))
50        })?
51    };
52    block_on_plugin(ctx, CalculatorService::with_config(calculator_config))
53}
54
55fn parse_tool_configs(config: &Value) -> Result<HashMap<String, ToolConfig>, cordis::CordisError> {
56    if config.is_null() {
57        return Ok(HashMap::new());
58    }
59    let Some(obj) = config.as_object() else {
60        return Err(cordis::CordisError::Configuration(
61            "invalid Tools config: expected object, {\"tools\": <map>}, empty, or null".into(),
62        ));
63    };
64    if obj.is_empty() {
65        return Ok(HashMap::new());
66    }
67    let map_value = if let Some(tools) = obj.get("tools") {
68        if tools.is_null() {
69            return Ok(HashMap::new());
70        }
71        tools.clone()
72    } else {
73        let mut stripped = config.clone();
74        if let Some(map) = stripped.as_object_mut() {
75            map.remove("mcps_dir");
76        }
77        stripped
78    };
79    if map_value.is_null() || map_value.as_object().is_some_and(serde_json::Map::is_empty) {
80        return Ok(HashMap::new());
81    }
82    serde_json::from_value(map_value).map_err(|error| {
83        cordis::CordisError::Configuration(format!("invalid Tools config: {error}"))
84    })
85}
86
87fn factory_tools(
88    ctx: &Arc<cordis::Context>,
89    config: &Value,
90) -> Result<cordis::FiberId, cordis::CordisError> {
91    let map = parse_tool_configs(config)?;
92    let mut tool_registry = ToolRegistry::with_config(&map);
93
94    tool_registry.register(Arc::new(Calculator));
95
96    #[cfg(feature = "search-tools")]
97    {
98        tool_registry.register(Arc::new(crate::search::WebSearch::new()));
99        tool_registry.register(Arc::new(crate::web_scrape::WebScrape::new()));
100    }
101
102    #[cfg(feature = "postgres")]
103    register_connector_tools(ctx, &mut tool_registry);
104
105    #[cfg(feature = "mcp")]
106    register_mcp_bridge_tools(config, &mut tool_registry);
107
108    tracing::info!(
109        "Tool registry initialized with {} tools",
110        tool_registry.enabled_tool_names().len()
111    );
112
113    let static_reg = Arc::new(tool_registry);
114
115    #[cfg(any(feature = "postgres", test))]
116    let tools = Tools::with_runtime(static_reg, runtime_registry(ctx));
117    #[cfg(not(any(feature = "postgres", test)))]
118    let tools = Tools::new(static_reg);
119
120    block_on_plugin(ctx, tools)
121}
122
123#[cfg(feature = "postgres")]
124fn register_connector_tools(ctx: &Arc<cordis::Context>, tool_registry: &mut ToolRegistry) {
125    match (
126        ares_store::MasterKey::from_env(),
127        ctx.get::<ares_store::PostgresClient>(),
128    ) {
129        (Some(master_key), Some(pg)) => {
130            crate::connectors::register_prebuilt_connector_tools(
131                tool_registry,
132                pg.pool.clone(),
133                master_key,
134            );
135        }
136        (Some(_), None) => {
137            tracing::warn!("PostgresClient missing; pre-built connector tools are not registered");
138        }
139        (None, _) => {
140            tracing::warn!(
141                "FLEET_SECRETS_KEY is not set; pre-built connector tools are not registered"
142            );
143        }
144    }
145}
146
147#[cfg(feature = "mcp")]
148fn register_mcp_bridge_tools(config: &Value, tool_registry: &mut ToolRegistry) {
149    let mcps_dir = config
150        .get("mcps_dir")
151        .and_then(Value::as_str)
152        .unwrap_or("config/mcps");
153    let Ok(mcp_reg) = ares_mcp::McpRegistry::from_dir(mcps_dir) else {
154        return;
155    };
156    for client_name in mcp_reg.client_names() {
157        if mcp_reg.get_client(&client_name).is_some() {
158            crate::mcp_bridge::register_mcp_tools(tool_registry, &client_name);
159        }
160    }
161}
162
163#[cfg(any(feature = "postgres", test))]
164fn runtime_registry(ctx: &Arc<cordis::Context>) -> Option<Arc<crate::RuntimeToolRegistry>> {
165    #[cfg(feature = "postgres")]
166    {
167        let pg = ctx.get::<ares_store::PostgresClient>()?;
168        let runtime_tool_registry = crate::RuntimeToolRegistry::new(pg.pool.clone());
169        if let Err(e) = tokio::task::block_in_place(|| {
170            tokio::runtime::Handle::current().block_on(runtime_tool_registry.reload())
171        }) {
172            tracing::warn!("Failed to preload runtime tools on startup: {}", e);
173        }
174        Some(Arc::new(runtime_tool_registry))
175    }
176    #[cfg(not(feature = "postgres"))]
177    {
178        let _ = ctx;
179        None
180    }
181}
182
183#[cfg(test)]
184mod tests {
185    use super::*;
186    use serde_json::json;
187
188    #[test]
189    fn register_plugins_registers_exactly_calculator_and_tools() {
190        let reg = cordis::PluginRegistry::new();
191        register_plugins(&reg);
192        let mut names = reg.names();
193        names.sort();
194        assert_eq!(
195            names,
196            vec!["CalculatorService".to_string(), "Tools".to_string()]
197        );
198    }
199
200    #[test]
201    fn parse_tool_configs_null_and_empty() {
202        assert!(parse_tool_configs(&Value::Null).unwrap().is_empty());
203        assert!(parse_tool_configs(&json!({})).unwrap().is_empty());
204    }
205
206    #[test]
207    fn parse_tool_configs_accepts_tools_wrapper() {
208        let map = parse_tool_configs(&json!({
209            "tools": {
210                "calculator": { "enabled": true }
211            },
212            "mcps_dir": "config/mcps"
213        }))
214        .unwrap();
215        assert!(map.get("calculator").is_some());
216        assert!(map.get("calculator").unwrap().enabled);
217    }
218
219    #[test]
220    fn parse_tool_configs_flat_map_strips_mcps_dir() {
221        let map = parse_tool_configs(&json!({
222            "calculator": { "enabled": false },
223            "mcps_dir": "elsewhere"
224        }))
225        .unwrap();
226        assert_eq!(map.len(), 1);
227        assert!(!map.get("calculator").unwrap().enabled);
228    }
229}