embacle_mcp/tools/
multiplex.rs1use async_trait::async_trait;
8use serde_json::{json, Value};
9
10use dravr_tronc::mcp::schema::{Tool, ToolResponse};
11use dravr_tronc::{McpTool, ToolContext};
12
13use crate::runner::{parse_runner_type, valid_provider_names, ALL_PROVIDERS};
14use crate::state::{ServerState, SharedState};
15
16pub struct GetMultiplexProvider;
18
19#[async_trait]
20impl McpTool<ServerState> for GetMultiplexProvider {
21 fn definition(&self) -> Tool {
22 Tool {
23 name: "get_multiplex_provider".to_owned(),
24 description: "Get the list of providers configured for multiplex prompt dispatch"
25 .to_owned(),
26 input_schema: json!({
27 "type": "object",
28 "properties": {}
29 }),
30 annotations: None,
31 }
32 }
33
34 async fn execute(
35 &self,
36 state: &SharedState,
37 _ctx: &ToolContext,
38 _arguments: Value,
39 ) -> ToolResponse {
40 let providers: Vec<String> = state
41 .multiplex_providers()
42 .await
43 .iter()
44 .map(ToString::to_string)
45 .collect();
46 let all: Vec<String> = ALL_PROVIDERS.iter().map(ToString::to_string).collect();
47
48 ToolResponse::text(
49 json!({
50 "multiplex_providers": providers,
51 "available_providers": all
52 })
53 .to_string(),
54 )
55 }
56}
57
58pub struct SetMultiplexProvider;
60
61#[async_trait]
62impl McpTool<ServerState> for SetMultiplexProvider {
63 fn definition(&self) -> Tool {
64 let provider_names: Vec<String> = ALL_PROVIDERS.iter().map(ToString::to_string).collect();
65
66 Tool {
67 name: "set_multiplex_provider".to_owned(),
68 description:
69 "Set providers for multiplex mode — prompts will fan out to all listed providers"
70 .to_owned(),
71 input_schema: json!({
72 "type": "object",
73 "properties": {
74 "providers": {
75 "type": "array",
76 "description": "List of provider names to multiplex to",
77 "items": {
78 "type": "string",
79 "enum": provider_names
80 }
81 }
82 },
83 "required": ["providers"]
84 }),
85 annotations: None,
86 }
87 }
88
89 async fn execute(
90 &self,
91 state: &SharedState,
92 _ctx: &ToolContext,
93 arguments: Value,
94 ) -> ToolResponse {
95 let Some(provider_strs) = arguments.get("providers").and_then(Value::as_array) else {
96 return ToolResponse::error("Missing 'providers' array argument".to_owned());
97 };
98
99 let mut providers = Vec::with_capacity(provider_strs.len());
100 for val in provider_strs {
101 let Some(name) = val.as_str() else {
102 return ToolResponse::error(format!("Provider must be a string, got: {val}"));
103 };
104 match parse_runner_type(name) {
105 Some(p) => providers.push(p),
106 None => {
107 return ToolResponse::error(format!(
108 "Unknown provider: {name}. Valid: {}",
109 valid_provider_names()
110 ));
111 }
112 }
113 }
114
115 let result_names: Vec<String> = providers.iter().map(ToString::to_string).collect();
116 state.set_multiplex_providers(providers).await;
117
118 ToolResponse::text(
119 json!({
120 "multiplex_providers": result_names,
121 "status": "configured"
122 })
123 .to_string(),
124 )
125 }
126}