Skip to main content

embacle_mcp/tools/
prompt.rs

1// ABOUTME: MCP tool that dispatches chat prompts to the active embacle LLM provider
2// ABOUTME: Supports single-provider and multiplex modes for concurrent multi-provider queries
3//
4// SPDX-License-Identifier: Apache-2.0
5// Copyright (c) 2026 dravr.ai
6
7use std::sync::Arc;
8
9use async_trait::async_trait;
10use embacle::types::{ChatMessage, ChatRequest, MessageRole};
11use serde_json::{json, Value};
12
13use dravr_tronc::mcp::schema::{Tool, ToolResponse};
14use dravr_tronc::{McpTool, ToolContext};
15
16use crate::runner::multiplex::MultiplexEngine;
17use crate::state::{ServerState, SharedState};
18
19/// Dispatches a chat prompt to the active provider or fans out via multiplex
20pub struct Prompt;
21
22#[async_trait]
23impl McpTool<ServerState> for Prompt {
24    fn definition(&self) -> Tool {
25        Tool {
26            name: "prompt".to_owned(),
27            description:
28                "Send a chat prompt to the active LLM provider, or multiplex to all configured providers"
29                    .to_owned(),
30            input_schema: json!({
31                "type": "object",
32                "properties": {
33                    "messages": {
34                        "type": "array",
35                        "description": "Chat messages to send to the provider",
36                        "items": {
37                            "type": "object",
38                            "properties": {
39                                "role": {
40                                    "type": "string",
41                                    "enum": ["system", "user", "assistant"]
42                                },
43                                "content": {
44                                    "type": "string"
45                                },
46                                "images": {
47                                    "type": "array",
48                                    "description": "Optional images attached to the message (user role only)",
49                                    "items": {
50                                        "type": "object",
51                                        "properties": {
52                                            "data": {
53                                                "type": "string",
54                                                "description": "Base64-encoded image data"
55                                            },
56                                            "mime_type": {
57                                                "type": "string",
58                                                "description": "MIME type (image/png, image/jpeg, image/webp, image/gif)"
59                                            }
60                                        },
61                                        "required": ["data", "mime_type"]
62                                    }
63                                }
64                            },
65                            "required": ["role", "content"]
66                        }
67                    },
68                    "multiplex": {
69                        "type": "boolean",
70                        "description": "If true, send to all multiplex providers instead of the active one",
71                        "default": false
72                    }
73                },
74                "required": ["messages"]
75            }),
76            annotations: None,
77        }
78    }
79
80    async fn execute(
81        &self,
82        state: &SharedState,
83        _ctx: &ToolContext,
84        arguments: Value,
85    ) -> ToolResponse {
86        let messages = match parse_messages(&arguments) {
87            Ok(msgs) => msgs,
88            Err(e) => return ToolResponse::error(e),
89        };
90
91        let multiplex = arguments
92            .get("multiplex")
93            .and_then(Value::as_bool)
94            .unwrap_or(false);
95
96        if multiplex {
97            execute_multiplex(state, &messages).await
98        } else {
99            execute_single(state, &messages).await
100        }
101    }
102}
103
104/// Execute a prompt against the single active provider
105async fn execute_single(state: &SharedState, messages: &[ChatMessage]) -> ToolResponse {
106    let provider = state.active_provider().await;
107    let runner = match state.get_runner(provider).await {
108        Ok(r) => r,
109        Err(e) => {
110            return ToolResponse::error(format!("Failed to create runner: {e}"));
111        }
112    };
113    let model = state.active_model().await;
114
115    let mut request = ChatRequest::new(messages.to_vec());
116    if let Some(m) = model {
117        request = request.with_model(m);
118    }
119
120    match runner.complete(&request).await {
121        Ok(response) => match serde_json::to_string_pretty(&response) {
122            Ok(json) => ToolResponse::text(json),
123            Err(e) => ToolResponse::error(format!("Response serialization failed: {e}")),
124        },
125        Err(e) => ToolResponse::error(format!("Completion error: {e}")),
126    }
127}
128
129/// Execute a prompt against all configured multiplex providers
130async fn execute_multiplex(state: &SharedState, messages: &[ChatMessage]) -> ToolResponse {
131    let providers = state.multiplex_providers().await;
132
133    if providers.is_empty() {
134        return ToolResponse::error(
135            "No multiplex providers configured. Use set_multiplex_provider first.".to_owned(),
136        );
137    }
138
139    let engine = MultiplexEngine::new(Arc::clone(state));
140    match engine.execute(messages, &providers).await {
141        Ok(result) => match serde_json::to_string_pretty(&result) {
142            Ok(json) => ToolResponse::text(json),
143            Err(e) => ToolResponse::error(format!("Result serialization failed: {e}")),
144        },
145        Err(e) => ToolResponse::error(format!("Multiplex error: {e}")),
146    }
147}
148
149/// Parse image objects from a message's "images" array
150fn parse_images(msg: &Value, index: usize) -> Result<Option<Vec<embacle::ImagePart>>, String> {
151    let Some(arr) = msg.get("images").and_then(Value::as_array) else {
152        return Ok(None);
153    };
154
155    if arr.is_empty() {
156        return Ok(None);
157    }
158
159    let mut images = Vec::with_capacity(arr.len());
160    for (j, img_val) in arr.iter().enumerate() {
161        let data = img_val
162            .get("data")
163            .and_then(Value::as_str)
164            .ok_or_else(|| format!("Message {index}, image {j}: missing 'data'"))?;
165        let mime_type = img_val
166            .get("mime_type")
167            .and_then(Value::as_str)
168            .ok_or_else(|| format!("Message {index}, image {j}: missing 'mime_type'"))?;
169
170        let part = embacle::ImagePart::new(data, mime_type)
171            .map_err(|e| format!("Message {index}, image {j}: {e}"))?;
172        images.push(part);
173    }
174
175    Ok(Some(images))
176}
177
178/// Parse chat messages from the MCP tool arguments JSON
179fn parse_messages(arguments: &Value) -> Result<Vec<ChatMessage>, String> {
180    let arr = arguments
181        .get("messages")
182        .and_then(Value::as_array)
183        .ok_or_else(|| "Missing or invalid 'messages' array".to_owned())?;
184
185    let mut messages = Vec::with_capacity(arr.len());
186    for (i, msg) in arr.iter().enumerate() {
187        let role_str = msg
188            .get("role")
189            .and_then(Value::as_str)
190            .ok_or_else(|| format!("Message {i}: missing 'role'"))?;
191
192        let content = msg
193            .get("content")
194            .and_then(Value::as_str)
195            .ok_or_else(|| format!("Message {i}: missing 'content'"))?;
196
197        let role = match role_str {
198            "system" => MessageRole::System,
199            "user" => MessageRole::User,
200            "assistant" => MessageRole::Assistant,
201            other => return Err(format!("Message {i}: invalid role '{other}'")),
202        };
203
204        let images = parse_images(msg, i)?;
205        let mut message = ChatMessage::new(role, content);
206        message.images = images;
207        messages.push(message);
208    }
209
210    if messages.is_empty() {
211        return Err("Messages array must not be empty".to_owned());
212    }
213
214    Ok(messages)
215}
216
217#[cfg(test)]
218mod tests {
219    use super::*;
220
221    #[test]
222    fn parse_valid_messages() {
223        let args = json!({
224            "messages": [
225                {"role": "system", "content": "You are helpful."},
226                {"role": "user", "content": "Hello!"}
227            ]
228        });
229        let msgs = parse_messages(&args).expect("should parse"); // Safe: test assertion
230        assert_eq!(msgs.len(), 2);
231        assert_eq!(msgs[0].role, MessageRole::System);
232        assert_eq!(msgs[1].content, "Hello!");
233    }
234
235    #[test]
236    fn parse_empty_messages_rejected() {
237        let args = json!({"messages": []});
238        assert!(parse_messages(&args).is_err());
239    }
240
241    #[test]
242    fn parse_missing_role_rejected() {
243        let args = json!({"messages": [{"content": "hi"}]});
244        assert!(parse_messages(&args).is_err());
245    }
246
247    #[test]
248    fn parse_invalid_role_rejected() {
249        let args = json!({"messages": [{"role": "bot", "content": "hi"}]});
250        let err = parse_messages(&args).unwrap_err();
251        assert!(err.contains("invalid role"));
252    }
253
254    #[test]
255    fn parse_messages_with_images() {
256        let args = json!({
257            "messages": [{
258                "role": "user",
259                "content": "Describe this",
260                "images": [{
261                    "data": "aGVsbG8=",
262                    "mime_type": "image/png"
263                }]
264            }]
265        });
266        let msgs = parse_messages(&args).expect("should parse"); // Safe: test assertion
267        assert_eq!(msgs.len(), 1);
268        let images = msgs[0].images.as_ref().expect("images present"); // Safe: test assertion
269        assert_eq!(images.len(), 1);
270        assert_eq!(images[0].mime_type, "image/png");
271        assert_eq!(images[0].data, "aGVsbG8=");
272    }
273
274    #[test]
275    fn parse_messages_without_images() {
276        let args = json!({
277            "messages": [{"role": "user", "content": "Hello!"}]
278        });
279        let msgs = parse_messages(&args).expect("should parse"); // Safe: test assertion
280        assert!(msgs[0].images.is_none());
281    }
282
283    #[test]
284    fn parse_messages_invalid_mime_type() {
285        let args = json!({
286            "messages": [{
287                "role": "user",
288                "content": "Describe",
289                "images": [{"data": "abc", "mime_type": "image/bmp"}]
290            }]
291        });
292        let err = parse_messages(&args).unwrap_err();
293        assert!(err.contains("image/bmp"));
294    }
295}