1use 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
22pub 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(®);
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}