Skip to main content

prism_mcp_rs/core/
prompt.rs

1//! Prompt system for MCP servers
2//!
3//! This module provides the abstraction for implementing and managing prompts in MCP servers.
4//! Prompts are templates that can be used to generate messages for language models.
5
6use 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/// Trait for implementing prompt handlers
17#[async_trait]
18pub trait PromptHandler: Send + Sync {
19    /// Generate prompt messages with the given arguments
20    ///
21    /// # Arguments
22    /// * `arguments` - Prompt arguments as key-value pairs
23    ///
24    /// # Returns
25    /// Result containing the generated prompt messages or an error
26    async fn get(&self, arguments: HashMap<String, Value>) -> McpResult<PromptResult>;
27}
28
29/// A registered prompt with its handler
30pub struct Prompt {
31    /// Information about the prompt
32    pub info: PromptInfo,
33    /// Handler that implements the prompt's functionality
34    pub handler: Box<dyn PromptHandler>,
35    /// Whether the prompt is currently enabled
36    pub enabled: bool,
37}
38
39impl Prompt {
40    /// Create a new prompt with the given information and handler
41    ///
42    /// # Arguments
43    /// * `info` - Information about the prompt
44    /// * `handler` - Implementation of the prompt's functionality
45    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    /// Enable the prompt
57    pub fn enable(&mut self) {
58        self.enabled = true;
59    }
60
61    /// Disable the prompt
62    pub fn disable(&mut self) {
63        self.enabled = false;
64    }
65
66    /// Check if the prompt is enabled
67    pub fn is_enabled(&self) -> bool {
68        self.enabled
69    }
70
71    /// Execute the prompt if it's enabled
72    ///
73    /// # Arguments
74    /// * `arguments` - Prompt arguments as key-value pairs
75    ///
76    /// # Returns
77    /// Result containing the prompt result or an error
78    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        // Validate required arguments
87        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    /// Create a system message
113    pub fn system<S: Into<String>>(content: S) -> Self {
114        Self {
115            role: Role::User, // Note: 2025-11-25 only has User and Assistant roles
116            content: Content::text(content.into()),
117        }
118    }
119
120    /// Create a user message
121    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    /// Create an assistant message
129    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    /// Create a message with custom role
137    pub fn with_role(role: Role, content: Content) -> Self {
138        Self { role, content }
139    }
140}
141
142// Common prompt implementations
143
144/// Simple greeting prompt
145pub 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
166/// Code review prompt
167pub 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
206/// SQL query generation prompt
207pub 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
247/// Builder for creating prompts with fluent API
248pub 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    /// Create a new prompt builder with the given name
258    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    /// Set the prompt description
269    pub fn description<S: Into<String>>(mut self, description: S) -> Self {
270        self.description = Some(description.into());
271        self
272    }
273
274    /// Set the prompt title (for UI display)
275    pub fn title<S: Into<String>>(mut self, title: S) -> Self {
276        self.title = Some(title.into());
277        self
278    }
279
280    /// Set prompt icons (for UI display)
281    pub fn icons(mut self, icons: Vec<Icon>) -> Self {
282        self.icons = Some(icons);
283        self
284    }
285
286    /// Add a single prompt icon (for UI display)
287    pub fn icon(mut self, icon: Icon) -> Self {
288        self.icons.get_or_insert_with(Vec::new).push(icon);
289        self
290    }
291
292    /// Add a required argument
293    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    /// Add an optional argument
304    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    /// Build the prompt with the given handler
315    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
336/// Utility for creating prompt arguments
337pub 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
346/// Utility for creating optional prompt arguments
347pub 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        // Test missing required argument
441        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}