1use 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
19pub 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
104async 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
129async 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
149fn 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
178fn 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"); 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"); assert_eq!(msgs.len(), 1);
268 let images = msgs[0].images.as_ref().expect("images present"); 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"); 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}