Skip to main content

embacle_mcp/tools/
multiplex.rs

1// ABOUTME: MCP tools for configuring multiplex providers that receive fan-out prompts
2// ABOUTME: Manages the list of providers used when prompt dispatch runs in multiplex mode
3//
4// SPDX-License-Identifier: Apache-2.0
5// Copyright (c) 2026 dravr.ai
6
7use 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
16/// Returns the list of providers configured for multiplex dispatch
17pub 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
58/// Sets the list of providers used when multiplexing prompts
59pub 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}