1use async_trait::async_trait;
7use serde_json::Value;
8use std::collections::HashMap;
9
10use crate::core::error::{McpError, McpResult};
11use crate::protocol::types::{
12 Content, GetPromptResult as PromptResult, Icon, Prompt as PromptInfo, PromptArgument,
13 PromptMessage, Role,
14};
15
16#[async_trait]
18pub trait PromptHandler: Send + Sync {
19 async fn get(&self, arguments: HashMap<String, Value>) -> McpResult<PromptResult>;
27}
28
29pub struct Prompt {
31 pub info: PromptInfo,
33 pub handler: Box<dyn PromptHandler>,
35 pub enabled: bool,
37}
38
39impl Prompt {
40 pub fn new<H>(info: PromptInfo, handler: H) -> Self
46 where
47 H: PromptHandler + 'static,
48 {
49 Self {
50 info,
51 handler: Box::new(handler),
52 enabled: true,
53 }
54 }
55
56 pub fn enable(&mut self) {
58 self.enabled = true;
59 }
60
61 pub fn disable(&mut self) {
63 self.enabled = false;
64 }
65
66 pub fn is_enabled(&self) -> bool {
68 self.enabled
69 }
70
71 pub async fn get(&self, arguments: HashMap<String, Value>) -> McpResult<PromptResult> {
79 if !self.enabled {
80 return Err(McpError::validation(format!(
81 "Prompt '{}' is disabled",
82 self.info.name
83 )));
84 }
85
86 if let Some(ref args) = self.info.arguments {
88 for arg in args {
89 if arg.required.unwrap_or(false) && !arguments.contains_key(&arg.name) {
90 return Err(McpError::validation(format!(
91 "Required argument '{}' missing for prompt '{}'",
92 arg.name, self.info.name
93 )));
94 }
95 }
96 }
97
98 self.handler.get(arguments).await
99 }
100}
101
102impl std::fmt::Debug for Prompt {
103 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
104 f.debug_struct("Prompt")
105 .field("info", &self.info)
106 .field("enabled", &self.enabled)
107 .finish()
108 }
109}
110
111impl PromptMessage {
112 pub fn system<S: Into<String>>(content: S) -> Self {
114 Self {
115 role: Role::User, content: Content::text(content.into()),
117 }
118 }
119
120 pub fn user<S: Into<String>>(content: S) -> Self {
122 Self {
123 role: Role::User,
124 content: Content::text(content.into()),
125 }
126 }
127
128 pub fn assistant<S: Into<String>>(content: S) -> Self {
130 Self {
131 role: Role::Assistant,
132 content: Content::text(content.into()),
133 }
134 }
135
136 pub fn with_role(role: Role, content: Content) -> Self {
138 Self { role, content }
139 }
140}
141
142pub struct GreetingPrompt;
146
147#[async_trait]
148impl PromptHandler for GreetingPrompt {
149 async fn get(&self, arguments: HashMap<String, Value>) -> McpResult<PromptResult> {
150 let name = arguments
151 .get("name")
152 .and_then(|v| v.as_str())
153 .unwrap_or("World");
154
155 Ok(PromptResult {
156 description: Some("A friendly greeting".to_string()),
157 messages: vec![
158 PromptMessage::system("You are a friendly assistant."),
159 PromptMessage::user(format!("Hello, {name}!")),
160 ],
161 meta: None,
162 })
163 }
164}
165
166pub struct CodeReviewPrompt;
168
169#[async_trait]
170impl PromptHandler for CodeReviewPrompt {
171 async fn get(&self, arguments: HashMap<String, Value>) -> McpResult<PromptResult> {
172 let code = arguments
173 .get("code")
174 .and_then(|v| v.as_str())
175 .ok_or_else(|| McpError::validation("Missing 'code' argument"))?;
176
177 let language = arguments
178 .get("language")
179 .and_then(|v| v.as_str())
180 .unwrap_or("unknown");
181
182 let focus = arguments
183 .get("focus")
184 .and_then(|v| v.as_str())
185 .unwrap_or("general");
186
187 let system_prompt = format!(
188 "You are an expert code reviewer. Focus on {focus} aspects of the code. \
189 Provide constructive feedback and suggestions for improvement."
190 );
191
192 let user_prompt =
193 format!("Please review this {language} code:\n\n```{language}\n{code}\n```");
194
195 Ok(PromptResult {
196 description: Some("Code review prompt".to_string()),
197 messages: vec![
198 PromptMessage::system(system_prompt),
199 PromptMessage::user(user_prompt),
200 ],
201 meta: None,
202 })
203 }
204}
205
206pub struct SqlQueryPrompt;
208
209#[async_trait]
210impl PromptHandler for SqlQueryPrompt {
211 async fn get(&self, arguments: HashMap<String, Value>) -> McpResult<PromptResult> {
212 let request = arguments
213 .get("request")
214 .and_then(|v| v.as_str())
215 .ok_or_else(|| McpError::validation("Missing 'request' argument"))?;
216
217 let schema = arguments
218 .get("schema")
219 .and_then(|v| v.as_str())
220 .unwrap_or("No schema provided");
221
222 let dialect = arguments
223 .get("dialect")
224 .and_then(|v| v.as_str())
225 .unwrap_or("PostgreSQL");
226
227 let system_prompt = format!(
228 "You are an expert SQL developer. Generate efficient and safe {dialect} queries. \
229 Always use proper escaping and avoid SQL injection vulnerabilities."
230 );
231
232 let user_prompt = format!(
233 "Database Schema:\n{schema}\n\nRequest: {request}\n\nPlease generate a {dialect} query for this request."
234 );
235
236 Ok(PromptResult {
237 description: Some("SQL query generation prompt".to_string()),
238 messages: vec![
239 PromptMessage::system(system_prompt),
240 PromptMessage::user(user_prompt),
241 ],
242 meta: None,
243 })
244 }
245}
246
247pub struct PromptBuilder {
249 name: String,
250 description: Option<String>,
251 arguments: Vec<PromptArgument>,
252 title: Option<String>,
253 icons: Option<Vec<Icon>>,
254}
255
256impl PromptBuilder {
257 pub fn new<S: Into<String>>(name: S) -> Self {
259 Self {
260 name: name.into(),
261 description: None,
262 arguments: Vec::new(),
263 title: None,
264 icons: None,
265 }
266 }
267
268 pub fn description<S: Into<String>>(mut self, description: S) -> Self {
270 self.description = Some(description.into());
271 self
272 }
273
274 pub fn title<S: Into<String>>(mut self, title: S) -> Self {
276 self.title = Some(title.into());
277 self
278 }
279
280 pub fn icons(mut self, icons: Vec<Icon>) -> Self {
282 self.icons = Some(icons);
283 self
284 }
285
286 pub fn icon(mut self, icon: Icon) -> Self {
288 self.icons.get_or_insert_with(Vec::new).push(icon);
289 self
290 }
291
292 pub fn required_arg<S: Into<String>>(mut self, name: S, description: Option<S>) -> Self {
294 self.arguments.push(PromptArgument {
295 name: name.into(),
296 description: description.map(|d| d.into()),
297 required: Some(true),
298 title: None,
299 });
300 self
301 }
302
303 pub fn optional_arg<S: Into<String>>(mut self, name: S, description: Option<S>) -> Self {
305 self.arguments.push(PromptArgument {
306 name: name.into(),
307 description: description.map(|d| d.into()),
308 required: Some(false),
309 title: None,
310 });
311 self
312 }
313
314 pub fn build<H>(self, handler: H) -> Prompt
316 where
317 H: PromptHandler + 'static,
318 {
319 let info = PromptInfo {
320 name: self.name,
321 description: self.description,
322 arguments: if self.arguments.is_empty() {
323 None
324 } else {
325 Some(self.arguments)
326 },
327 icons: self.icons,
328 title: self.title,
329 meta: None,
330 };
331
332 Prompt::new(info, handler)
333 }
334}
335
336pub fn required_arg<S: Into<String>>(name: S, description: Option<S>) -> PromptArgument {
338 PromptArgument {
339 name: name.into(),
340 description: description.map(|d| d.into()),
341 required: Some(true),
342 title: None,
343 }
344}
345
346pub fn optional_arg<S: Into<String>>(name: S, description: Option<S>) -> PromptArgument {
348 PromptArgument {
349 name: name.into(),
350 description: description.map(|d| d.into()),
351 required: Some(false),
352 title: None,
353 }
354}
355
356#[cfg(test)]
357mod tests {
358 use super::*;
359 use serde_json::json;
360
361 #[tokio::test]
362 async fn test_greeting_prompt() {
363 let prompt = GreetingPrompt;
364 let mut args = HashMap::new();
365 args.insert("name".to_string(), json!("Alice"));
366
367 let result = prompt.get(args).await.unwrap();
368 assert_eq!(result.messages.len(), 2);
369 assert_eq!(result.messages[0].role, Role::User);
370 assert_eq!(result.messages[1].role, Role::User);
371
372 match &result.messages[1].content {
373 Content::Text { text, .. } => assert!(text.contains("Alice")),
374 _ => panic!("Expected text content"),
375 }
376 }
377
378 #[tokio::test]
379 async fn test_code_review_prompt() {
380 let prompt = CodeReviewPrompt;
381 let mut args = HashMap::new();
382 args.insert(
383 "code".to_string(),
384 json!("function hello() { console.log('Hello'); }"),
385 );
386 args.insert("language".to_string(), json!("javascript"));
387 args.insert("focus".to_string(), json!("performance"));
388
389 let result = prompt.get(args).await.unwrap();
390 assert_eq!(result.messages.len(), 2);
391
392 match &result.messages[1].content {
393 Content::Text { text, .. } => {
394 assert!(text.contains("javascript"));
395 assert!(text.contains("console.log"));
396 }
397 _ => panic!("Expected text content"),
398 }
399 }
400
401 #[test]
402 fn test_prompt_creation() {
403 let info = PromptInfo {
404 name: "test_prompt".to_string(),
405 description: Some("Test prompt".to_string()),
406 arguments: Some(vec![PromptArgument {
407 name: "arg1".to_string(),
408 description: Some("First argument".to_string()),
409 required: Some(true),
410 title: None,
411 }]),
412 icons: None,
413 title: None,
414 meta: None,
415 };
416
417 let prompt = Prompt::new(info.clone(), GreetingPrompt);
418 assert_eq!(prompt.info, info);
419 assert!(prompt.is_enabled());
420 }
421
422 #[tokio::test]
423 async fn test_prompt_validation() {
424 let info = PromptInfo {
425 name: "test_prompt".to_string(),
426 description: None,
427 arguments: Some(vec![PromptArgument {
428 name: "required_arg".to_string(),
429 description: None,
430 required: Some(true),
431 title: None,
432 }]),
433 icons: None,
434 title: None,
435 meta: None,
436 };
437
438 let prompt = Prompt::new(info, GreetingPrompt);
439
440 let result = prompt.get(HashMap::new()).await;
442 assert!(result.is_err());
443 match result.unwrap_err() {
444 McpError::Validation(msg) => assert!(msg.contains("required_arg")),
445 _ => panic!("Expected validation error"),
446 }
447 }
448
449 #[test]
450 fn test_prompt_builder() {
451 let prompt = PromptBuilder::new("test")
452 .description("A test prompt")
453 .title("Test Prompt Title")
454 .icon(Icon {
455 src: "https://example.com/prompt-icon.svg".to_string(),
456 mime_type: Some("image/svg+xml".to_string()),
457 sizes: None,
458 theme: None,
459 })
460 .required_arg("input", Some("Input text"))
461 .optional_arg("format", Some("Output format"))
462 .build(GreetingPrompt);
463
464 assert_eq!(prompt.info.name, "test");
465 assert_eq!(prompt.info.description, Some("A test prompt".to_string()));
466 assert_eq!(prompt.info.title, Some("Test Prompt Title".to_string()));
467 assert_eq!(prompt.info.icons.as_ref().map(Vec::len), Some(1));
468
469 let args = prompt.info.arguments.unwrap();
470 assert_eq!(args.len(), 2);
471 assert_eq!(args[0].name, "input");
472 assert_eq!(args[0].required, Some(true));
473 assert_eq!(args[1].name, "format");
474 assert_eq!(args[1].required, Some(false));
475 }
476
477 #[test]
478 fn test_prompt_message_creation() {
479 let system_msg = PromptMessage::system("You are a helpful assistant");
480 assert_eq!(system_msg.role, Role::User);
481
482 let user_msg = PromptMessage::user("Hello!");
483 assert_eq!(user_msg.role, Role::User);
484
485 let assistant_msg = PromptMessage::assistant("Hi there!");
486 assert_eq!(assistant_msg.role, Role::Assistant);
487 }
488
489 #[test]
490 fn test_prompt_content_creation() {
491 let text_content = Content::text("Hello, world!");
492 match text_content {
493 Content::Text { text, .. } => {
494 assert_eq!(text, "Hello, world!");
495 }
496 _ => panic!("Expected text content"),
497 }
498
499 let image_content = Content::image("base64data", "image/png");
500 match image_content {
501 Content::Image {
502 data, mime_type, ..
503 } => {
504 assert_eq!(data, "base64data");
505 assert_eq!(mime_type, "image/png");
506 }
507 _ => panic!("Expected image content"),
508 }
509 }
510}