diff --git a/_context/wiki/config.md b/_context/wiki/config.md index 304c9a3c..b1dd5928 100644 --- a/_context/wiki/config.md +++ b/_context/wiki/config.md @@ -156,8 +156,9 @@ RuntimePluginConfigDocument cpex: CpexConfig ``` -Supported: `cmf.tool_pre_invoke`, `cmf.tool_post_invoke` only. -Rejected: routing-based selection, plugin dirs, global policies, other hook types. +Supported: `cmf.tool_pre_invoke`, `cmf.tool_post_invoke`, `cmf.prompt_pre_fetch`, `cmf.prompt_post_fetch` only. +Rejected: routing-based selection, plugin dirs, global policies, resource and LLM hooks, plugin conditions. +Config validation and `CmfPluginFactory` registration must agree on that list: a hook accepted by validation but not registered leaves the plugin loaded and silently inert. Reload watcher: 10-minute interval. Invalid reload → runtime marked failed. ### Tool Call Hook Behavior @@ -168,6 +169,22 @@ After the upstream backend returns, the post hook can leave the result unchanged Plugin execution must not poison shared gateway state. A plugin denial becomes an MCP error. Soft plugin errors are logged. Unsupported plugin configuration fails validation before the runtime is accepted. +### Prompt Fetch Hook Behavior + +For `get_prompt`, the pre hook runs after backend routing, so the plugin sees the backend-local prompt name and the owning backend separately rather than the gateway-prefixed identifier. It can leave the arguments unchanged, replace them, or deny the fetch before the backend renders anything. + +The post hook receives the rendered prompt as one CMF message per rendered MCP message, each carrying its role and its content block: text, image, audio, embedded resource, or resource link. A plugin can inspect or rewrite any of them, so a policy can act on a file interpolated into a prompt rather than only on the surrounding text. + +Writing plugin edits back follows three rules: + +- A message the plugin left unchanged is returned exactly as the backend sent it, so annotations, `_meta`, and binary resource blobs survive untouched. +- A message the plugin changed is rebuilt from CMF. CMF does not model MCP annotations or `_meta`, so an edited message loses them. +- Edits that cannot be applied faithfully fail the call rather than falling back to the backend's original. A changed message count, anything other than exactly one prompt result in the payload, a role MCP prompts cannot express, or a resource whose text the plugin removed all return an error. Silently restoring the backend's content would undo a redaction. + +MCP prompt results carry no error flag, so a plugin setting `is_error` on the CMF prompt result is rejecting the prompt rather than describing it. The gateway turns that into an MCP error carrying the plugin's `error_message`, and the rendered content never reaches the client. This differs from tools, where `is_error` is a field on `CallToolResult` and is forwarded as a successful response. + +Binary resource blobs reach plugins by URI and MIME type but not by content: CMF stores decoded bytes while MCP sends base64. A plugin can deny such a message; editing one fails the write-back. + ### Demo Plugin Workflow The optional `test-plugins` feature compiles demo factories from the `cpex-plugins-rs` repository. Redis configuration activates factories already present in the binary; it never loads new Rust code into a running process. diff --git a/_context/wiki/project.md b/_context/wiki/project.md index 9f45e2fa..4764a0b7 100644 --- a/_context/wiki/project.md +++ b/_context/wiki/project.md @@ -25,7 +25,7 @@ flowchart LR direction TB MW["Middleware stack\nvirtual host · JWT · session · user config"] RT["MCP Routing\nfan-out · prefix namespace\nlist merge · capability merge"] - PL["Plugin hooks\ncmf.tool_pre_invoke\ncmf.tool_post_invoke"] + PL["Plugin hooks\ncmf.tool_pre_invoke\ncmf.tool_post_invoke\ncmf.prompt_pre_fetch\ncmf.prompt_post_fetch"] MW --> RT --> PL end diff --git a/_context/wiki/routing.md b/_context/wiki/routing.md index 4a681a85..3ec9576a 100644 --- a/_context/wiki/routing.md +++ b/_context/wiki/routing.md @@ -128,7 +128,7 @@ If RMCP rejects the delete, local state is untouched. | `call_tool` | Targeted | Resolves alias → single/multi-backend fallback. Runs pre/post plugin hooks. Forwards downstream cancellation to backend. Tracks backend progress tokens: RMCP assigns a new token per backend request; the gateway maps each backend token to the downstream token. Request enqueue and mapping publication are serialized against progress lookup so an immediate backend notification cannot overtake registration. When the notification matches an in-flight token, the gateway restores the downstream token and forwards it to the client. | | `read_resource` | Targeted | Single-backend: URI unchanged. Multi-backend: strips prefix. | | `subscribe` / `unsubscribe` | Targeted | Same resource-URI routing; forwards/stops resource-update notifications. | -| `get_prompt` | Targeted | Single-backend: name unchanged. Multi-backend: strips prefix. | +| `get_prompt` | Targeted | Single-backend: name unchanged. Multi-backend: strips prefix. Runs pre/post prompt hooks around the backend call: the pre hook may rewrite arguments or deny, the post hook may rewrite or reject the rendered messages. | | `complete` | Targeted | Routes on prompt name or resource URI inside `ref`. | | `ping` | Local | Returns success; no backend fanout. | | `DELETE` | Session | RMCP handles first; on success `session_id_layer` removes local session + backend transports. | diff --git a/_context/wiki/testing.md b/_context/wiki/testing.md index 8d16a434..6ea10ac4 100644 --- a/_context/wiki/testing.md +++ b/_context/wiki/testing.md @@ -30,7 +30,7 @@ Protocol tests and fixtures should target MCP `2026-07-28`, use `server/discover | `gateway_list_tools.rs` | List fanout, prefixing, and merged output. | | `gateway_prompts.rs` | Prompt listing and prefixed `get_prompt` routing. | | `gateway_resource_templates.rs` | Template fanout with prefixed names and URI templates, plus `read_resource` round-trips. | -| `gateway_plugins.rs` | CPEX pre/post tool hooks around `call_tool` and stream events. | +| `gateway_plugins.rs` | CPEX pre/post tool hooks around `call_tool` and stream events, and prompt hooks around `get_prompt`. | These run in `cargo nextest run` with no Docker dependencies. diff --git a/crates/contextforge-data-plane-cpex/src/cmf.rs b/crates/contextforge-data-plane-cpex/src/cmf.rs index b4de6f6e..eeb74c97 100644 --- a/crates/contextforge-data-plane-cpex/src/cmf.rs +++ b/crates/contextforge-data-plane-cpex/src/cmf.rs @@ -1,5 +1,13 @@ -use cpex::cpex_core::cmf::{ContentPart, Message, MessagePayload, Role, ToolCall, ToolResult}; -use rmcp::model::{CallToolRequestParams, CallToolResult, ContentBlock}; +use std::collections::HashMap; + +use cpex::cpex_core::cmf::{ + AudioSource, ContentPart, ImageSource, Message, MessagePayload, PromptRequest, PromptResult, + Resource as CmfResource, ResourceReference, ResourceType, Role, ToolCall, ToolResult, +}; +use rmcp::model::{ + CallToolRequestParams, CallToolResult, ContentBlock, GetPromptRequestParams, GetPromptResult, PromptMessage, + Resource as McpResource, ResourceContents, Role as McpRole, +}; use serde_json::{Map, Value}; pub(crate) fn tool_call_payload( @@ -110,10 +118,670 @@ fn raw_error_tool_result(value: Value) -> CallToolResult { } } +pub(crate) fn prompt_request_payload( + request: &GetPromptRequestParams, + prompt_name: &str, + backend_name: &str, + prompt_request_id: &str, +) -> MessagePayload { + MessagePayload { + message: Message { + schema_version: "2.0".to_owned(), + role: Role::User, + content: vec![ContentPart::PromptRequest { + content: PromptRequest { + prompt_request_id: prompt_request_id.to_owned(), + name: prompt_name.to_owned(), + arguments: request.arguments.clone().map(HashMap::from_iter).unwrap_or_default(), + server_id: Some(backend_name.to_owned()), + }, + }], + channel: None, + }, + } +} + +pub(crate) fn prompt_request_arguments( + payload: &MessagePayload, + prompt_name: &str, + backend_name: &str, + prompt_request_id: &str, +) -> Option> { + let requests = payload.message.get_prompt_requests(); + let [request] = requests.as_slice() else { return None }; + if request.name != prompt_name + || request.prompt_request_id != prompt_request_id + || request.server_id.as_deref() != Some(backend_name) + { + return None; + } + Some(request.arguments.clone().into_iter().collect::>()) +} + +pub(crate) fn prompt_result_payload( + response: &GetPromptResult, + prompt_name: &str, + prompt_request_id: &str, +) -> MessagePayload { + let messages = + response.messages.iter().map(|message| cmf_prompt_message(message, prompt_request_id)).collect::>(); + + MessagePayload { + message: Message { + schema_version: "2.0".to_owned(), + role: Role::Assistant, + content: vec![ContentPart::PromptResult { + content: PromptResult { + prompt_request_id: prompt_request_id.to_owned(), + prompt_name: prompt_name.to_owned(), + messages, + content: None, + is_error: false, + error_message: None, + }, + }], + channel: None, + }, + } +} + +fn prompt_result(payload: &MessagePayload) -> Option<&PromptResult> { + let results = payload.message.get_prompt_results(); + let [result] = results.as_slice() else { return None }; + Some(*result) +} + +pub(crate) fn prompt_result_rejection(payload: &MessagePayload) -> Option { + let result = prompt_result(payload)?; + result + .is_error + .then(|| result.error_message.clone().unwrap_or_else(|| "Plugin rejected the rendered prompt".to_owned())) +} + +// `None` means refuse: falling back to the backend's original would undo a plugin's redaction. +pub(crate) fn prompt_result_response( + mut original: GetPromptResult, + payload: &MessagePayload, + prompt_name: &str, + prompt_request_id: &str, +) -> Option { + let result = prompt_result(payload)?; + if result.prompt_name != prompt_name + || result.prompt_request_id != prompt_request_id + || result.content.is_some() + || result.error_message.is_some() + { + return None; + } + if result.messages.len() != original.messages.len() { + return None; + } + + for (message, edited) in original.messages.iter_mut().zip(&result.messages) { + let projected = cmf_prompt_message(message, prompt_request_id); + if serde_json::to_value(&projected).ok()? == serde_json::to_value(edited).ok()? { + continue; + } + + let rebuilt = mcp_prompt_message(edited)?; + if serde_json::to_value(cmf_prompt_message(&rebuilt, prompt_request_id)).ok()? + != serde_json::to_value(edited).ok()? + { + return None; + } + *message = rebuilt; + } + + Some(original) +} + +fn cmf_prompt_message(message: &PromptMessage, prompt_request_id: &str) -> Message { + Message { + schema_version: "2.0".to_owned(), + role: match message.role { + McpRole::Assistant => Role::Assistant, + McpRole::User => Role::User, + }, + content: cmf_content_part(&message.content, prompt_request_id).into_iter().collect(), + channel: None, + } +} + +fn cmf_content_part(block: &ContentBlock, prompt_request_id: &str) -> Option { + let part = match block { + ContentBlock::Text(text) => ContentPart::Text { text: text.text.clone() }, + ContentBlock::Image(image) => ContentPart::Image { + content: ImageSource { + source_type: "base64".to_owned(), + data: image.data.clone(), + media_type: Some(image.mime_type.clone()), + }, + }, + ContentBlock::Audio(audio) => ContentPart::Audio { + content: AudioSource { + source_type: "base64".to_owned(), + data: audio.data.clone(), + media_type: Some(audio.mime_type.clone()), + duration_ms: None, + }, + }, + ContentBlock::Resource(resource) => { + let (uri, mime_type, content) = match &resource.resource { + ResourceContents::TextResourceContents { uri, mime_type, text, .. } => { + (uri.clone(), mime_type.clone(), Some(text.clone())) + }, + ResourceContents::BlobResourceContents { uri, mime_type, .. } => (uri.clone(), mime_type.clone(), None), + _ => return None, + }; + ContentPart::Resource { + content: CmfResource { + resource_request_id: prompt_request_id.to_owned(), + uri, + name: None, + description: None, + resource_type: ResourceType::Uri, + content, + blob: None, + mime_type, + size_bytes: None, + annotations: HashMap::new(), + version: None, + }, + } + }, + ContentBlock::ResourceLink(link) => ContentPart::ResourceRef { + content: ResourceReference { + resource_request_id: prompt_request_id.to_owned(), + uri: link.uri.clone(), + name: Some(link.name.clone()), + resource_type: ResourceType::Uri, + range_start: None, + range_end: None, + selector: None, + }, + }, + _ => return None, + }; + + Some(part) +} + +// MCP inlines image and audio bytes as base64, so a CMF source CMF can express but MCP cannot — +// a URL reference — has to be refused rather than written into a field that means something else. +fn inline_media_data<'a>(source_type: &str, data: &'a str) -> Option<&'a str> { + (source_type == "base64").then_some(data) +} + +fn mcp_prompt_message(message: &Message) -> Option { + let role = match message.role { + Role::Assistant => McpRole::Assistant, + Role::User => McpRole::User, + _ => return None, + }; + + let [part] = message.content.as_slice() else { return None }; + let content = match part { + ContentPart::Text { text } => ContentBlock::text(text.clone()), + ContentPart::Image { content } => { + ContentBlock::image(inline_media_data(&content.source_type, &content.data)?, content.media_type.clone()?) + }, + ContentPart::Audio { content } => { + ContentBlock::audio(inline_media_data(&content.source_type, &content.data)?, content.media_type.clone()?) + }, + ContentPart::Resource { content } => ContentBlock::resource(ResourceContents::TextResourceContents { + uri: content.uri.clone(), + mime_type: content.mime_type.clone(), + text: content.content.clone()?, + meta: None, + }), + ContentPart::ResourceRef { content } => { + ContentBlock::ResourceLink(McpResource::new(content.uri.clone(), content.name.clone()?)) + }, + _ => return None, + }; + + Some(PromptMessage::new(role, content)) +} + #[cfg(test)] mod tests { use super::*; + fn text_prompt() -> GetPromptResult { + GetPromptResult::new(vec![PromptMessage::new_text(McpRole::User, "review of weather")]) + } + + fn prompt_result_mut(payload: &mut MessagePayload) -> &mut PromptResult { + payload + .message + .content + .iter_mut() + .find_map(|part| match part { + ContentPart::PromptResult { content } => Some(content), + _ => None, + }) + .expect("payload carries a prompt result") + } + + fn edited_messages(payload: &mut MessagePayload) -> &mut Vec { + &mut prompt_result_mut(payload).messages + } + + #[test] + fn prompt_result_response_rejects_added_message() { + let original = text_prompt(); + let mut payload = prompt_result_payload(&original, "review", "prompt-1"); + let extra = edited_messages(&mut payload).first().cloned().expect("one message"); + edited_messages(&mut payload).push(extra); + + assert!(prompt_result_response(original, &payload, "review", "prompt-1").is_none()); + } + + #[test] + fn prompt_result_response_rejects_extra_prompt_result() { + let original = text_prompt(); + let mut payload = prompt_result_payload(&original, "review", "prompt-1"); + let duplicate = payload.message.content[0].clone(); + payload.message.content.push(duplicate); + + assert!(prompt_result_response(original, &payload, "review", "prompt-1").is_none()); + } + + #[test] + fn prompt_result_rejection_reports_the_plugin_error_message() { + let original = text_prompt(); + let mut payload = prompt_result_payload(&original, "review", "prompt-1"); + let result = prompt_result_mut(&mut payload); + result.is_error = true; + result.error_message = Some("blocked by policy".to_owned()); + + assert_eq!(Some("blocked by policy".to_owned()), prompt_result_rejection(&payload)); + } + + #[test] + fn prompt_result_rejection_falls_back_when_the_plugin_gives_no_message() { + let original = text_prompt(); + let mut payload = prompt_result_payload(&original, "review", "prompt-1"); + prompt_result_mut(&mut payload).is_error = true; + + assert_eq!(Some("Plugin rejected the rendered prompt".to_owned()), prompt_result_rejection(&payload)); + } + + #[test] + fn prompt_result_rejection_is_absent_for_a_normal_result() { + let original = text_prompt(); + let payload = prompt_result_payload(&original, "review", "prompt-1"); + + assert_eq!(None, prompt_result_rejection(&payload)); + } + + fn review_payload() -> MessagePayload { + let request = GetPromptRequestParams::new("review") + .with_arguments(Map::from_iter([("topic".to_owned(), Value::from("weather"))])); + prompt_request_payload(&request, "review", "backend-a", "prompt-1") + } + + fn prompt_request_mut(payload: &mut MessagePayload) -> &mut PromptRequest { + payload + .message + .content + .iter_mut() + .find_map(|part| match part { + ContentPart::PromptRequest { content } => Some(content), + _ => None, + }) + .expect("payload carries a prompt request") + } + + #[test] + fn prompt_request_arguments_accepts_an_argument_edit() { + let mut payload = review_payload(); + prompt_request_mut(&mut payload).arguments.insert("topic".to_owned(), Value::from("rain")); + + let arguments = prompt_request_arguments(&payload, "review", "backend-a", "prompt-1"); + + assert_eq!(Some(&Value::from("rain")), arguments.as_ref().and_then(|args| args.get("topic"))); + } + + #[test] + fn prompt_request_arguments_rejects_a_renamed_prompt() { + let mut payload = review_payload(); + "other".clone_into(&mut prompt_request_mut(&mut payload).name); + + assert!(prompt_request_arguments(&payload, "review", "backend-a", "prompt-1").is_none()); + } + + #[test] + fn prompt_request_arguments_rejects_a_rerouted_backend() { + let mut payload = review_payload(); + prompt_request_mut(&mut payload).server_id = Some("backend-b".to_owned()); + + assert!(prompt_request_arguments(&payload, "review", "backend-a", "prompt-1").is_none()); + } + + #[test] + fn prompt_request_arguments_rejects_a_recorrelated_request() { + let mut payload = review_payload(); + "prompt-2".clone_into(&mut prompt_request_mut(&mut payload).prompt_request_id); + + assert!(prompt_request_arguments(&payload, "review", "backend-a", "prompt-1").is_none()); + } + + #[test] + fn prompt_request_arguments_rejects_extra_prompt_requests() { + let mut payload = review_payload(); + let duplicate = payload.message.content[0].clone(); + payload.message.content.push(duplicate); + + assert!(prompt_request_arguments(&payload, "review", "backend-a", "prompt-1").is_none()); + } + + #[test] + fn prompt_result_response_rejects_envelope_content_edit() { + let original = text_prompt(); + let mut payload = prompt_result_payload(&original, "review", "prompt-1"); + prompt_result_mut(&mut payload).content = Some("[REDACTED]".to_owned()); + + assert!(prompt_result_response(original, &payload, "review", "prompt-1").is_none()); + } + + #[test] + fn prompt_result_response_rejects_renamed_prompt() { + let original = text_prompt(); + let mut payload = prompt_result_payload(&original, "review", "prompt-1"); + prompt_result_mut(&mut payload).prompt_name = "other".to_owned(); + + assert!(prompt_result_response(original, &payload, "review", "prompt-1").is_none()); + } + + #[test] + fn prompt_result_response_rejects_recorrelated_result() { + let original = text_prompt(); + let mut payload = prompt_result_payload(&original, "review", "prompt-1"); + prompt_result_mut(&mut payload).prompt_request_id = "prompt-2".to_owned(); + + assert!(prompt_result_response(original, &payload, "review", "prompt-1").is_none()); + } + + #[test] + fn prompt_result_response_rejects_error_message_without_error_flag() { + let original = text_prompt(); + let mut payload = prompt_result_payload(&original, "review", "prompt-1"); + prompt_result_mut(&mut payload).error_message = Some("blocked".to_owned()); + + assert!(prompt_result_response(original, &payload, "review", "prompt-1").is_none()); + } + + fn resource_prompt() -> GetPromptResult { + GetPromptResult::new(vec![PromptMessage::new( + McpRole::User, + ContentBlock::resource(ResourceContents::text("token=secret", "file:///app.env")), + )]) + } + + #[test] + fn prompt_result_response_rejects_resource_type_edit() { + let original = resource_prompt(); + let mut payload = prompt_result_payload(&original, "review", "prompt-1"); + let ContentPart::Resource { content } = &mut edited_messages(&mut payload)[0].content[0] else { + panic!("expected a resource part"); + }; + content.resource_type = ResourceType::Database; + + assert!(prompt_result_response(original, &payload, "review", "prompt-1").is_none()); + } + + #[test] + fn prompt_result_response_rejects_dropped_resource_metadata() { + let original = resource_prompt(); + let mut payload = prompt_result_payload(&original, "review", "prompt-1"); + let ContentPart::Resource { content } = &mut edited_messages(&mut payload)[0].content[0] else { + panic!("expected a resource part"); + }; + content.description = Some("annotated by policy".to_owned()); + + assert!(prompt_result_response(original, &payload, "review", "prompt-1").is_none()); + } + + fn media_prompt(content: ContentBlock) -> GetPromptResult { + GetPromptResult::new(vec![PromptMessage::new(McpRole::User, content)]) + } + + #[test] + fn prompt_result_response_round_trips_an_image_edit() { + let original = media_prompt(ContentBlock::image("aW1hZ2U=", "image/png")); + let mut payload = prompt_result_payload(&original, "review", "prompt-1"); + let ContentPart::Image { content } = &mut edited_messages(&mut payload)[0].content[0] else { + panic!("expected an image part"); + }; + content.data = "cmVkYWN0ZWQ=".to_owned(); + + let result = prompt_result_response(original, &payload, "review", "prompt-1").expect("image edit applies"); + + let ContentBlock::Image(image) = &result.messages[0].content else { panic!("expected an image") }; + assert_eq!("cmVkYWN0ZWQ=", image.data); + assert_eq!("image/png", image.mime_type); + } + + #[test] + fn prompt_result_response_round_trips_an_audio_edit() { + let original = media_prompt(ContentBlock::audio("YXVkaW8=", "audio/mp3")); + let mut payload = prompt_result_payload(&original, "review", "prompt-1"); + let ContentPart::Audio { content } = &mut edited_messages(&mut payload)[0].content[0] else { + panic!("audio reaches the plugin as a CMF audio part"); + }; + content.data = "cmVkYWN0ZWQ=".to_owned(); + + let result = prompt_result_response(original, &payload, "review", "prompt-1").expect("audio edit applies"); + + let ContentBlock::Audio(audio) = &result.messages[0].content else { panic!("expected audio") }; + assert_eq!("cmVkYWN0ZWQ=", audio.data); + assert_eq!("audio/mp3", audio.mime_type); + } + + #[test] + fn prompt_result_response_rejects_url_sourced_audio() { + let original = media_prompt(ContentBlock::audio("YXVkaW8=", "audio/mp3")); + let mut payload = prompt_result_payload(&original, "review", "prompt-1"); + let ContentPart::Audio { content } = &mut edited_messages(&mut payload)[0].content[0] else { + panic!("expected an audio part"); + }; + "url".clone_into(&mut content.source_type); + content.data = "https://example.invalid/clip.mp3".to_owned(); + + assert!(prompt_result_response(original, &payload, "review", "prompt-1").is_none()); + } + + #[test] + fn prompt_result_response_rejects_audio_without_media_type() { + let original = media_prompt(ContentBlock::audio("YXVkaW8=", "audio/mp3")); + let mut payload = prompt_result_payload(&original, "review", "prompt-1"); + let ContentPart::Audio { content } = &mut edited_messages(&mut payload)[0].content[0] else { + panic!("expected an audio part"); + }; + content.media_type = None; + + assert!(prompt_result_response(original, &payload, "review", "prompt-1").is_none()); + } + + #[test] + fn prompt_result_response_round_trips_a_resource_link_edit() { + let original = media_prompt(ContentBlock::ResourceLink(McpResource::new("file:///app.env", "app-env"))); + let mut payload = prompt_result_payload(&original, "review", "prompt-1"); + let ContentPart::ResourceRef { content } = &mut edited_messages(&mut payload)[0].content[0] else { + panic!("expected a resource reference part"); + }; + content.name = Some("redacted-env".to_owned()); + + let result = prompt_result_response(original, &payload, "review", "prompt-1").expect("link edit applies"); + + let ContentBlock::ResourceLink(link) = &result.messages[0].content else { panic!("expected a link") }; + assert_eq!("redacted-env", link.name); + assert_eq!("file:///app.env", link.uri); + } + + #[test] + fn prompt_result_response_rejects_resource_with_removed_text() { + let original = resource_prompt(); + let mut payload = prompt_result_payload(&original, "review", "prompt-1"); + let ContentPart::Resource { content } = &mut edited_messages(&mut payload)[0].content[0] else { + panic!("expected a resource part"); + }; + content.content = None; + + assert!(prompt_result_response(original, &payload, "review", "prompt-1").is_none()); + } + + #[test] + fn prompt_result_response_rejects_multiple_content_parts() { + let original = text_prompt(); + let mut payload = prompt_result_payload(&original, "review", "prompt-1"); + edited_messages(&mut payload)[0].content.push(ContentPart::Text { text: "extra".to_owned() }); + + assert!(prompt_result_response(original, &payload, "review", "prompt-1").is_none()); + } + + #[test] + fn prompt_result_response_rejects_a_cmf_only_content_part() { + let original = text_prompt(); + let mut payload = prompt_result_payload(&original, "review", "prompt-1"); + edited_messages(&mut payload)[0].content = vec![ContentPart::Thinking { text: "reasoning".to_owned() }]; + + assert!(prompt_result_response(original, &payload, "review", "prompt-1").is_none()); + } + + #[test] + fn prompt_result_response_rejects_a_payload_without_a_prompt_result() { + let original = text_prompt(); + let mut payload = prompt_result_payload(&original, "review", "prompt-1"); + payload.message.content.clear(); + + assert!(prompt_result_response(original, &payload, "review", "prompt-1").is_none()); + } + + #[test] + fn prompt_request_arguments_rejects_a_payload_without_a_prompt_request() { + let mut payload = review_payload(); + payload.message.content.clear(); + + assert!(prompt_request_arguments(&payload, "review", "backend-a", "prompt-1").is_none()); + } + + #[test] + fn prompt_result_response_rejects_url_sourced_image() { + let original = + GetPromptResult::new(vec![PromptMessage::new(McpRole::User, ContentBlock::image("aW1hZ2U=", "image/png"))]); + let mut payload = prompt_result_payload(&original, "review", "prompt-1"); + let ContentPart::Image { content } = &mut edited_messages(&mut payload)[0].content[0] else { + panic!("expected an image part"); + }; + "url".clone_into(&mut content.source_type); + content.data = "https://example.invalid/image.png".to_owned(); + + assert!(prompt_result_response(original, &payload, "review", "prompt-1").is_none()); + } + + #[test] + fn prompt_result_response_rejects_image_without_media_type() { + let original = + GetPromptResult::new(vec![PromptMessage::new(McpRole::User, ContentBlock::image("aW1hZ2U=", "image/png"))]); + let mut payload = prompt_result_payload(&original, "review", "prompt-1"); + let ContentPart::Image { content } = &mut edited_messages(&mut payload)[0].content[0] else { + panic!("expected an image part"); + }; + content.media_type = None; + + assert!(prompt_result_response(original, &payload, "review", "prompt-1").is_none()); + } + + #[test] + fn prompt_result_response_rejects_resource_link_without_name() { + let original = GetPromptResult::new(vec![PromptMessage::new( + McpRole::User, + ContentBlock::ResourceLink(McpResource::new("file:///app.env", "app-env")), + )]); + let mut payload = prompt_result_payload(&original, "review", "prompt-1"); + let ContentPart::ResourceRef { content } = &mut edited_messages(&mut payload)[0].content[0] else { + panic!("expected a resource reference part"); + }; + content.name = None; + + assert!(prompt_result_response(original, &payload, "review", "prompt-1").is_none()); + } + + #[test] + fn prompt_result_response_rejects_resource_link_range_edit() { + let original = GetPromptResult::new(vec![PromptMessage::new( + McpRole::User, + ContentBlock::ResourceLink(McpResource::new("file:///app.env", "app-env")), + )]); + let mut payload = prompt_result_payload(&original, "review", "prompt-1"); + let ContentPart::ResourceRef { content } = &mut edited_messages(&mut payload)[0].content[0] else { + panic!("expected a resource reference part"); + }; + content.range_start = Some(10); + + assert!(prompt_result_response(original, &payload, "review", "prompt-1").is_none()); + } + + #[test] + fn prompt_result_response_rejects_removed_message() { + let original = text_prompt(); + let mut payload = prompt_result_payload(&original, "review", "prompt-1"); + edited_messages(&mut payload).clear(); + + assert!(prompt_result_response(original, &payload, "review", "prompt-1").is_none()); + } + + #[test] + fn prompt_result_response_rejects_unmappable_role() { + let original = text_prompt(); + let mut payload = prompt_result_payload(&original, "review", "prompt-1"); + edited_messages(&mut payload)[0].role = Role::System; + + assert!(prompt_result_response(original, &payload, "review", "prompt-1").is_none()); + } + + #[test] + fn prompt_result_response_preserves_unmodified_messages() { + let original = text_prompt(); + let payload = prompt_result_payload(&original, "review", "prompt-1"); + + let result = prompt_result_response(original.clone(), &payload, "review", "prompt-1") + .expect("unmodified payload applies"); + + assert_eq!( + serde_json::to_value(&original).expect("original serializes"), + serde_json::to_value(&result).expect("result serializes") + ); + } + + #[test] + fn prompt_result_response_round_trips_embedded_resource() { + let original = GetPromptResult::new(vec![PromptMessage::new( + McpRole::User, + ContentBlock::resource(ResourceContents::text("token=secret", "file:///app.env")), + )]); + let mut payload = prompt_result_payload(&original, "review", "prompt-1"); + + let ContentPart::Resource { content } = &mut edited_messages(&mut payload)[0].content[0] else { + panic!("embedded resource reaches the plugin as a CMF resource part"); + }; + assert_eq!(Some("token=secret"), content.content.as_deref()); + content.content = Some("token=[REDACTED]".to_owned()); + + let result = prompt_result_response(original, &payload, "review", "prompt-1").expect("resource edit applies"); + + let ContentBlock::Resource(resource) = &result.messages[0].content else { + panic!("expected an embedded resource"); + }; + let ResourceContents::TextResourceContents { text, uri, .. } = &resource.resource else { + panic!("expected text resource contents"); + }; + assert_eq!("token=[REDACTED]", text); + assert_eq!("file:///app.env", uri); + } + #[test] fn tool_result_response_uses_cmf_error_flag_for_nested_mcp_result() { let original = CallToolResult::success(vec![ContentBlock::text("original")]); diff --git a/crates/contextforge-data-plane-cpex/src/factory.rs b/crates/contextforge-data-plane-cpex/src/factory.rs index a41c4f40..598d8a2e 100644 --- a/crates/contextforge-data-plane-cpex/src/factory.rs +++ b/crates/contextforge-data-plane-cpex/src/factory.rs @@ -1,4 +1,4 @@ -use std::{marker::PhantomData, sync::Arc}; +use std::sync::Arc; use cpex::cpex_core::{ cmf::CmfHook, @@ -10,13 +10,12 @@ use cpex::cpex_core::{ }; pub struct CmfPluginFactory

{ - build: fn(PluginConfig) -> P, - _plugin: PhantomData

, + build: Box P + Send + Sync>, } impl

CmfPluginFactory

{ - pub fn new(build: fn(PluginConfig) -> P) -> Self { - Self { build, _plugin: PhantomData } + pub fn new(build: impl Fn(PluginConfig) -> P + Send + Sync + 'static) -> Self { + Self { build: Box::new(build) } } } @@ -50,6 +49,8 @@ fn cmf_hook_name(hook: &str) -> Option<&'static str> { match hook { cmf_hook_names::TOOL_PRE_INVOKE => Some(cmf_hook_names::TOOL_PRE_INVOKE), cmf_hook_names::TOOL_POST_INVOKE => Some(cmf_hook_names::TOOL_POST_INVOKE), + cmf_hook_names::PROMPT_PRE_FETCH => Some(cmf_hook_names::PROMPT_PRE_FETCH), + cmf_hook_names::PROMPT_POST_FETCH => Some(cmf_hook_names::PROMPT_POST_FETCH), _ => None, } } diff --git a/crates/contextforge-data-plane-cpex/src/handle.rs b/crates/contextforge-data-plane-cpex/src/handle.rs index c8c4dfe2..eae5f182 100644 --- a/crates/contextforge-data-plane-cpex/src/handle.rs +++ b/crates/contextforge-data-plane-cpex/src/handle.rs @@ -13,7 +13,7 @@ use cpex::cpex_core::{ }; use rmcp::{ ErrorData, - model::{CallToolRequestParams, CallToolResult, ErrorCode}, + model::{CallToolRequestParams, CallToolResult, ErrorCode, GetPromptRequestParams, GetPromptResult}, serde::{Serialize, de::DeserializeOwned}, }; use tokio::task::JoinHandle; @@ -21,7 +21,7 @@ use tokio::task::JoinHandle; use crate::{ config::{LoadedRuntimePluginConfig, RedisRuntimePluginConfigStore, RuntimePluginConfigStore, cpex_config}, error::GatewayPluginRuntimeError, - hooks::{RuntimeHookError, RuntimeHookState, ToolPreCallResult}, + hooks::{PromptPreFetchResult, RuntimeHookError, RuntimeHookState, ToolPreCallResult}, runtime::GatewayPluginRuntime, }; @@ -40,7 +40,7 @@ pub struct GatewayPluginRuntimeHandle { runtime: Arc>, } -struct RegistryToolCallState { +struct RegistryCallState { runtime: Arc, state: Option, } @@ -253,20 +253,52 @@ impl GatewayPluginRuntimeHandle { let mut result = runtime.before_tool_call(request, tool_name, backend_name).await?; if runtime.has_post_hook() { let state = result.state.take(); - result.state = Some(Arc::new(RegistryToolCallState { runtime: Arc::clone(runtime), state })); + result.state = Some(Arc::new(RegistryCallState { runtime: Arc::clone(runtime), state })); } else { result.state = None; } Ok(result) } + pub async fn before_get_prompt( + &self, + request: &GetPromptRequestParams, + prompt_name: &str, + backend_name: &str, + ) -> Result { + let state = self.current(); + let RuntimeState::Active(runtime) = state.as_ref() else { + return Err(runtime_failed_error(state.as_ref())); + }; + let mut result = runtime.before_get_prompt(request, prompt_name, backend_name).await?; + if runtime.has_prompt_post_hook() { + let state = result.state.take(); + result.state = Some(Arc::new(RegistryCallState { runtime: Arc::clone(runtime), state })); + } else { + result.state = None; + } + Ok(result) + } + + pub async fn after_get_prompt( + &self, + prompt_name: &str, + response: GetPromptResult, + state: Option, + ) -> Result { + match state.and_then(|state| state.downcast::().ok()) { + Some(state) => state.runtime.after_get_prompt(prompt_name, response, state.state.clone()).await, + None => Ok(response), + } + } + pub async fn after_tool_call( &self, tool_name: &str, response: CallToolResult, state: Option, ) -> Result { - match state.and_then(|state| state.downcast::().ok()) { + match state.and_then(|state| state.downcast::().ok()) { Some(state) => state.runtime.after_tool_call(tool_name, response, state.state.clone()).await, None => Ok(response), } @@ -283,7 +315,7 @@ impl GatewayPluginRuntimeHandle { where T: Serialize + DeserializeOwned, { - match state.and_then(|state| state.downcast::().ok()) { + match state.and_then(|state| state.downcast::().ok()) { Some(state) => state.runtime.after_tool_event(tool_name, event, state.state.clone()).await, None => Ok(Some(event)), } @@ -329,13 +361,14 @@ mod tests { }; use crate::config::LoadedRuntimePluginConfig; - use crate::{CmfPluginFactory, ToolArgumentsUpdate}; + use crate::{CmfPluginFactory, PromptArgumentsUpdate, ToolArgumentsUpdate}; use super::*; const TEST_MISSING_CONTEXT_ERROR_CODE: i64 = -32003; const TEST_REWRITTEN_SUM_A: i64 = 10; const TEST_REWRITTEN_SUM_B: i64 = 20; + const TEST_REWRITTEN_PROMPT_TOPIC: &str = "rewritten-topic"; const TEST_SHUTDOWN_RETRY_COUNT: usize = 20; const TEST_SHUTDOWN_RETRY_INTERVAL: Duration = Duration::from_millis(10); const TEST_WATCHER_INTERVAL: Duration = Duration::from_millis(10); @@ -555,6 +588,15 @@ mod tests { ("b".to_owned(), json!(TEST_REWRITTEN_SUM_B)), ]); } + if let Some(ContentPart::PromptRequest { content }) = modified + .message + .content + .iter_mut() + .find(|part| matches!(part, ContentPart::PromptRequest { .. })) + { + content.arguments = + HashMap::from([("topic".to_owned(), json!(TEST_REWRITTEN_PROMPT_TOPIC))]); + } PluginResult::modify_payload(modified) }, PreBehavior::SetContext => { @@ -616,6 +658,11 @@ mod tests { .with_arguments(serde_json::Map::from_iter([("a".to_owned(), json!(a)), ("b".to_owned(), json!(b))])) } + fn review_request(topic: &str) -> GetPromptRequestParams { + GetPromptRequestParams::new("review") + .with_arguments(serde_json::Map::from_iter([("topic".to_owned(), json!(topic))])) + } + fn progress_event() -> ProgressNotificationParam { ProgressNotificationParam::new(ProgressToken(NumberOrString::String("stream-token".into())), 1.0) .with_message("step 1/2") @@ -714,6 +761,13 @@ mod tests { } } + #[tokio::test(flavor = "multi_thread", worker_threads = 1)] + async fn prompt_hooks_are_accepted_config() { + let plugin = Arc::new(TestPlugin::new("prompt", vec![cmf_hook_names::PROMPT_PRE_FETCH])); + // runtime_with_plugin initializes and expects success + runtime_with_plugin(&plugin, plugin_config(&[Arc::clone(&plugin)])).await; + } + #[tokio::test(flavor = "multi_thread", worker_threads = 1)] async fn runtime_config_loads_registered_factory_plugin() { let plugin = @@ -747,6 +801,59 @@ mod tests { assert!(matches!(result.arguments, ToolArgumentsUpdate::Replace(Some(_)))); } + #[tokio::test(flavor = "multi_thread", worker_threads = 1)] + async fn generic_cmf_factory_registers_prompt_only_plugin() { + let config = config_document(json!({ + "plugins": [{ + "name": "generic-prompt", + "kind": "generic", + "hooks": [cmf_hook_names::PROMPT_PRE_FETCH] + }] + })); + let mut runtime = CpexRuntimeRegistry::with_config_store(Arc::new(MemoryConfigStore::with_config(config))); + runtime + .register_factory("generic", Box::new(CmfPluginFactory::new(TestPlugin::rewrite_from_config))) + .expect("test factory registers"); + runtime.initialize().await.expect("runtime initializes"); + + let result = runtime + .handle() + .before_get_prompt(&review_request("weather"), "review", "backend") + .await + .expect("prompt pre hook runs"); + + assert!( + matches!(result.arguments, PromptArgumentsUpdate::Replace(Some(_))), + "the prompt hook must actually run, not merely be accepted by config validation" + ); + } + + #[tokio::test(flavor = "multi_thread", worker_threads = 1)] + async fn generic_cmf_factory_registers_mixed_tool_and_prompt_plugin() { + let config = config_document(json!({ + "plugins": [{ + "name": "generic-mixed", + "kind": "generic", + "hooks": [cmf_hook_names::TOOL_PRE_INVOKE, cmf_hook_names::PROMPT_PRE_FETCH] + }] + })); + let mut runtime = CpexRuntimeRegistry::with_config_store(Arc::new(MemoryConfigStore::with_config(config))); + runtime + .register_factory("generic", Box::new(CmfPluginFactory::new(TestPlugin::rewrite_from_config))) + .expect("test factory registers"); + runtime.initialize().await.expect("runtime initializes"); + + let tool = runtime.before_tool_call(&sum_request(1, 2), "sum", "backend").await.expect("tool pre hook runs"); + let prompt = runtime + .handle() + .before_get_prompt(&review_request("weather"), "review", "backend") + .await + .expect("prompt pre hook runs"); + + assert!(matches!(tool.arguments, ToolArgumentsUpdate::Replace(Some(_)))); + assert!(matches!(prompt.arguments, PromptArgumentsUpdate::Replace(Some(_)))); + } + #[tokio::test(flavor = "multi_thread", worker_threads = 1)] async fn runtime_reload_replaces_and_clears_current_runtime() { let plugin = diff --git a/crates/contextforge-data-plane-cpex/src/hooks.rs b/crates/contextforge-data-plane-cpex/src/hooks.rs index 0e6ce655..b6a9a63a 100644 --- a/crates/contextforge-data-plane-cpex/src/hooks.rs +++ b/crates/contextforge-data-plane-cpex/src/hooks.rs @@ -1,6 +1,6 @@ use std::{any::Any, sync::Arc}; -use rmcp::model::CallToolRequestParams; +use rmcp::model::{CallToolRequestParams, GetPromptRequestParams}; use serde_json::{Map, Value}; pub type RuntimeHookError = Box; @@ -31,3 +31,29 @@ impl ToolPreCallResult { Self { arguments: ToolArgumentsUpdate::Unchanged, state: None } } } + +#[derive(Debug)] +pub enum PromptArgumentsUpdate { + Unchanged, + Replace(Option>), +} + +impl PromptArgumentsUpdate { + pub fn apply_to_request(self, request: &mut GetPromptRequestParams, routed_prompt_name: &str) { + routed_prompt_name.clone_into(&mut request.name); + if let Self::Replace(arguments) = self { + request.arguments = arguments; + } + } +} + +pub struct PromptPreFetchResult { + pub arguments: PromptArgumentsUpdate, + pub state: Option, +} + +impl PromptPreFetchResult { + pub fn unchanged() -> Self { + Self { arguments: PromptArgumentsUpdate::Unchanged, state: None } + } +} diff --git a/crates/contextforge-data-plane-cpex/src/lib.rs b/crates/contextforge-data-plane-cpex/src/lib.rs index 6a57cb51..5d1be674 100644 --- a/crates/contextforge-data-plane-cpex/src/lib.rs +++ b/crates/contextforge-data-plane-cpex/src/lib.rs @@ -10,4 +10,7 @@ mod runtime; pub use error::GatewayPluginRuntimeError; pub use factory::CmfPluginFactory; pub use handle::{CpexRuntimeRegistry, GatewayPluginRuntimeHandle}; -pub use hooks::{RuntimeHookError, RuntimeHookState, ToolArgumentsUpdate, ToolPreCallResult}; +pub use hooks::{ + PromptArgumentsUpdate, PromptPreFetchResult, RuntimeHookError, RuntimeHookState, ToolArgumentsUpdate, + ToolPreCallResult, +}; diff --git a/crates/contextforge-data-plane-cpex/src/pipeline.rs b/crates/contextforge-data-plane-cpex/src/pipeline.rs index 5d6cdd97..9e4ecab0 100644 --- a/crates/contextforge-data-plane-cpex/src/pipeline.rs +++ b/crates/contextforge-data-plane-cpex/src/pipeline.rs @@ -2,14 +2,17 @@ use cpex::cpex_core::cmf::MessagePayload; use cpex::cpex_core::executor::PipelineResult; use rmcp::{ ErrorData, - model::{CallToolResult, ErrorCode}, + model::{CallToolResult, ErrorCode, GetPromptResult}, serde::de::DeserializeOwned, }; use tracing::warn; use crate::{ - ToolArgumentsUpdate, - cmf::{tool_call_arguments, tool_result_content, tool_result_response}, + PromptArgumentsUpdate, ToolArgumentsUpdate, + cmf::{ + prompt_request_arguments, prompt_result_rejection, prompt_result_response, tool_call_arguments, + tool_result_content, tool_result_response, + }, }; pub(crate) fn modified_message_payload(result: &PipelineResult) -> Option<&MessagePayload> { @@ -39,6 +42,33 @@ pub(crate) fn effective_pre_args( } } +pub(crate) fn effective_pre_prompt_args( + original_args: Option<&serde_json::Map>, + pre_result: &PipelineResult, + prompt_name: &str, + backend_name: &str, + prompt_request_id: &str, +) -> Result { + let Some(modified_payload) = modified_message_payload(pre_result) else { + return Ok(PromptArgumentsUpdate::Unchanged); + }; + + let Some(arguments) = prompt_request_arguments(modified_payload, prompt_name, backend_name, prompt_request_id) + else { + return Err(ErrorData { + code: ErrorCode::INVALID_PARAMS, + message: "Plugin returned a prompt request the gateway cannot apply".into(), + data: None, + }); + }; + + if original_args == Some(&arguments) || (original_args.is_none() && arguments.is_empty()) { + Ok(PromptArgumentsUpdate::Unchanged) + } else { + Ok(PromptArgumentsUpdate::Replace(Some(arguments))) + } +} + pub(crate) fn effective_post_result(original: CallToolResult, result: &PipelineResult) -> CallToolResult { match modified_message_payload(result) { Some(payload) => tool_result_response(original, payload), @@ -46,6 +76,27 @@ pub(crate) fn effective_post_result(original: CallToolResult, result: &PipelineR } } +pub(crate) fn effective_post_prompt_result( + original: GetPromptResult, + result: &PipelineResult, + prompt_name: &str, + prompt_request_id: &str, +) -> Result { + let Some(payload) = modified_message_payload(result) else { + return Ok(original); + }; + + if let Some(message) = prompt_result_rejection(payload) { + return Err(ErrorData { code: ErrorCode::INVALID_REQUEST, message: message.into(), data: None }); + } + + prompt_result_response(original, payload, prompt_name, prompt_request_id).ok_or_else(|| ErrorData { + code: ErrorCode::INTERNAL_ERROR, + message: "Plugin returned a prompt result the gateway cannot apply".into(), + data: None, + }) +} + pub(crate) fn effective_post_json(original: T, result: &PipelineResult) -> Result where T: DeserializeOwned, @@ -67,16 +118,16 @@ where }) } -pub(crate) fn plugin_denied_error(result: PipelineResult) -> ErrorData { +pub(crate) fn plugin_denied_error(subject: &str, result: PipelineResult) -> ErrorData { let code = result .violation .and_then(|violation| { - warn!("Plugin denied tool call: code={} plugin={:?}", violation.code, violation.plugin_name); + warn!("Plugin denied {subject}: code={} plugin={:?}", violation.code, violation.plugin_name); violation.proto_error_code.and_then(|code| i32::try_from(code).ok()).map(ErrorCode) }) .unwrap_or(ErrorCode::INVALID_REQUEST); - ErrorData { code, message: "Plugin denied tool call".into(), data: None } + ErrorData { code, message: format!("Plugin denied {subject}").into(), data: None } } pub(crate) fn log_pipeline_errors(hook: &'static str, result: &PipelineResult) { diff --git a/crates/contextforge-data-plane-cpex/src/runtime.rs b/crates/contextforge-data-plane-cpex/src/runtime.rs index 21dc39c5..36c2efb7 100644 --- a/crates/contextforge-data-plane-cpex/src/runtime.rs +++ b/crates/contextforge-data-plane-cpex/src/runtime.rs @@ -14,25 +14,39 @@ use cpex::cpex_core::{ }; use rmcp::{ ErrorData, - model::{CallToolRequestParams, CallToolResult}, + model::{CallToolRequestParams, CallToolResult, GetPromptRequestParams, GetPromptResult}, serde::{Serialize, de::DeserializeOwned}, }; use tokio::sync::Mutex; use crate::{ - cmf::{tool_call_payload, tool_json_result_payload, tool_result_payload}, + cmf::{ + prompt_request_payload, prompt_result_payload, tool_call_payload, tool_json_result_payload, tool_result_payload, + }, error::GatewayPluginRuntimeError, - hooks::{RuntimeHookState, ToolArgumentsUpdate, ToolPreCallResult}, + hooks::{PromptPreFetchResult, RuntimeHookState, ToolArgumentsUpdate, ToolPreCallResult}, pipeline::{ - effective_post_json, effective_post_result, effective_pre_args, log_pipeline_errors, plugin_denied_error, + effective_post_json, effective_post_prompt_result, effective_post_result, effective_pre_args, + effective_pre_prompt_args, log_pipeline_errors, plugin_denied_error, }, }; +#[derive(Default)] +struct HookPair { + pre: bool, + post: bool, +} + +#[derive(Default)] +struct HookPresence { + tool: HookPair, + prompt: HookPair, +} + #[derive(Default)] pub(crate) struct GatewayPluginRuntime { manager: PluginManager, - has_pre_hook: bool, - has_post_hook: bool, + hooks: HookPresence, } struct ToolCallState { @@ -42,10 +56,10 @@ struct ToolCallState { type SharedToolCallState = Mutex; -static TOOL_CALL_ID: AtomicU64 = AtomicU64::new(1); +static CORRELATION_ID: AtomicU64 = AtomicU64::new(1); fn next_tool_call_id() -> String { - format!("gateway-tool-call-{}", TOOL_CALL_ID.fetch_add(1, Ordering::Relaxed)) + format!("gateway-tool-call-{}", CORRELATION_ID.fetch_add(1, Ordering::Relaxed)) } fn new_tool_call_state() -> RuntimeHookState { @@ -55,9 +69,26 @@ fn new_tool_call_state() -> RuntimeHookState { })) } +fn next_prompt_request_id() -> String { + format!("gateway-prompt-request-{}", CORRELATION_ID.fetch_add(1, Ordering::Relaxed)) +} + +struct PromptCallState { + context_table: PluginContextTable, + prompt_request_id: String, +} + +fn new_prompt_call_state(context_table: PluginContextTable, prompt_request_id: String) -> RuntimeHookState { + Arc::new(PromptCallState { context_table, prompt_request_id }) +} + impl GatewayPluginRuntime { pub(crate) fn has_post_hook(&self) -> bool { - self.has_post_hook + self.hooks.tool.post + } + + pub(crate) fn has_prompt_post_hook(&self) -> bool { + self.hooks.prompt.post } pub(crate) async fn from_config( @@ -66,16 +97,20 @@ impl GatewayPluginRuntime { ) -> Result { validate_gateway_supported_config(&config)?; - let has_pre_hook = - config.plugins.iter().any(|plugin| plugin.hooks.iter().any(|hook| hook == cmf_hook_names::TOOL_PRE_INVOKE)); - let has_post_hook = config - .plugins - .iter() - .any(|plugin| plugin.hooks.iter().any(|hook| hook == cmf_hook_names::TOOL_POST_INVOKE)); + let hooks = HookPresence { + tool: HookPair { + pre: declares(&config, cmf_hook_names::TOOL_PRE_INVOKE), + post: declares(&config, cmf_hook_names::TOOL_POST_INVOKE), + }, + prompt: HookPair { + pre: declares(&config, cmf_hook_names::PROMPT_PRE_FETCH), + post: declares(&config, cmf_hook_names::PROMPT_POST_FETCH), + }, + }; let manager = PluginManager::from_config(config, factories) .map_err(|source| GatewayPluginRuntimeError::Configuration { hook: "config", source })?; manager.initialize().await.map_err(|source| GatewayPluginRuntimeError::Initialization { source })?; - Ok(Self { manager, has_pre_hook, has_post_hook }) + Ok(Self { manager, hooks }) } } @@ -93,6 +128,17 @@ impl Drop for GatewayPluginRuntime { } } +const SUPPORTED_HOOKS: [&str; 4] = [ + cmf_hook_names::TOOL_PRE_INVOKE, + cmf_hook_names::TOOL_POST_INVOKE, + cmf_hook_names::PROMPT_PRE_FETCH, + cmf_hook_names::PROMPT_POST_FETCH, +]; + +fn declares(config: &CpexConfig, hook_name: &str) -> bool { + config.plugins.iter().any(|plugin| plugin.hooks.iter().any(|hook| hook == hook_name)) +} + fn validate_gateway_supported_config(config: &CpexConfig) -> Result<(), GatewayPluginRuntimeError> { if config.routing_enabled() || config.plugin_settings.fail_on_plugin_error @@ -109,11 +155,7 @@ fn validate_gateway_supported_config(config: &CpexConfig) -> Result<(), GatewayP return Err(GatewayPluginRuntimeError::ConfigUnsupported); } - if plugin - .hooks - .iter() - .any(|hook| hook != cmf_hook_names::TOOL_PRE_INVOKE && hook != cmf_hook_names::TOOL_POST_INVOKE) - { + if plugin.hooks.iter().any(|hook| !SUPPORTED_HOOKS.contains(&hook.as_str())) { return Err(GatewayPluginRuntimeError::ConfigUnsupported); } } @@ -152,8 +194,8 @@ impl GatewayPluginRuntime { tool_name: &str, backend_name: &str, ) -> Result { - if !self.has_pre_hook { - let state = self.has_post_hook.then(new_tool_call_state); + if !self.hooks.tool.pre { + let state = self.hooks.tool.post.then(new_tool_call_state); return Ok(ToolPreCallResult { arguments: ToolArgumentsUpdate::Unchanged, state }); } @@ -161,7 +203,7 @@ impl GatewayPluginRuntime { let original_payload = tool_call_payload(request, tool_name, backend_name, &tool_call_id); let pre_result = self.invoke_tool_pre(original_payload).await; if pre_result.is_denied() { - return Err(plugin_denied_error(pre_result)); + return Err(plugin_denied_error("tool call", pre_result)); } let arguments = effective_pre_args(request.arguments.as_ref(), &pre_result)?; @@ -169,13 +211,94 @@ impl GatewayPluginRuntime { Ok(ToolPreCallResult { arguments, state: Some(Arc::new(state)) }) } + async fn invoke_prompt_pre(&self, payload: MessagePayload) -> PipelineResult { + let (result, background_tasks) = self + .manager + .invoke_named::(cmf_hook_names::PROMPT_PRE_FETCH, payload, Extensions::default(), None) + .await; + log_pipeline_errors(cmf_hook_names::PROMPT_PRE_FETCH, &result); + drop(background_tasks); + result + } + + async fn invoke_prompt_post( + &self, + payload: MessagePayload, + context_table: Option, + ) -> PipelineResult { + let (result, background_tasks) = self + .manager + .invoke_named::(cmf_hook_names::PROMPT_POST_FETCH, payload, Extensions::default(), context_table) + .await; + log_pipeline_errors(cmf_hook_names::PROMPT_POST_FETCH, &result); + drop(background_tasks); + result + } + + pub(crate) async fn before_get_prompt( + &self, + request: &GetPromptRequestParams, + prompt_name: &str, + backend_name: &str, + ) -> Result { + if !self.hooks.prompt.pre { + let mut result = PromptPreFetchResult::unchanged(); + result.state = self + .hooks + .prompt + .post + .then(|| new_prompt_call_state(PluginContextTable::default(), next_prompt_request_id())); + return Ok(result); + } + + let prompt_request_id = next_prompt_request_id(); + let payload = prompt_request_payload(request, prompt_name, backend_name, &prompt_request_id); + let pre_result = self.invoke_prompt_pre(payload).await; + if pre_result.is_denied() { + return Err(plugin_denied_error("prompt", pre_result)); + } + + let arguments = effective_pre_prompt_args( + request.arguments.as_ref(), + &pre_result, + prompt_name, + backend_name, + &prompt_request_id, + )?; + let state = + self.hooks.prompt.post.then(|| new_prompt_call_state(pre_result.context_table.clone(), prompt_request_id)); + Ok(PromptPreFetchResult { arguments, state }) + } + + pub(crate) async fn after_get_prompt( + &self, + prompt_name: &str, + response: GetPromptResult, + state: Option, + ) -> Result { + if !self.hooks.prompt.post { + return Ok(response); + } + + let state = state.and_then(|state| state.downcast::().ok()); + let Some(state) = state else { return Ok(response) }; + + let payload = prompt_result_payload(&response, prompt_name, &state.prompt_request_id); + let post_result = self.invoke_prompt_post(payload, Some(state.context_table.clone())).await; + if post_result.is_denied() { + return Err(plugin_denied_error("prompt", post_result)); + } + + effective_post_prompt_result(response, &post_result, prompt_name, &state.prompt_request_id) + } + pub(crate) async fn after_tool_call( &self, tool_name: &str, response: CallToolResult, state: Option, ) -> Result { - if !self.has_post_hook { + if !self.hooks.tool.post { return Ok(response); } @@ -190,7 +313,7 @@ impl GatewayPluginRuntime { ) .await; if post_result.is_denied() { - return Err(plugin_denied_error(post_result)); + return Err(plugin_denied_error("tool call", post_result)); } state.context_table = post_result.context_table.clone(); @@ -206,7 +329,7 @@ impl GatewayPluginRuntime { where T: Serialize + DeserializeOwned, { - if !self.has_post_hook { + if !self.hooks.tool.post { return Ok(Some(event)); } diff --git a/crates/contextforge-data-plane-lib/src/gateway/mcp_service/prompts.rs b/crates/contextforge-data-plane-lib/src/gateway/mcp_service/prompts.rs index 14c0654d..4962441e 100644 --- a/crates/contextforge-data-plane-lib/src/gateway/mcp_service/prompts.rs +++ b/crates/contextforge-data-plane-lib/src/gateway/mcp_service/prompts.rs @@ -1,3 +1,4 @@ +use contextforge_data_plane_cpex::PromptPreFetchResult; use rmcp::{ ErrorData, RoleServer, model::{GetPromptRequestParams, GetPromptResponse, ListPromptsResult, PaginatedRequestParams}, @@ -84,12 +85,22 @@ where ) .await?; + let pre_result = if let Some(plugin_runtime) = &mcp_service.plugin_runtime { + plugin_runtime.before_get_prompt(&request, &prompt_name, &service_name).await? + } else { + PromptPreFetchResult::unchanged() + }; let mut routed_request = request; - routed_request.name = prompt_name; + pre_result.arguments.apply_to_request(&mut routed_request, &prompt_name); let response = service .get_prompt(routed_request) .await .map_err(|error| backend_forward_error("get_prompt", &service_name, &error))?; info!("get_prompt: backend {service_name} returned {} messages", response.messages.len()); + let response = if let Some(plugin_runtime) = &mcp_service.plugin_runtime { + plugin_runtime.after_get_prompt(&prompt_name, response, pre_result.state).await? + } else { + response + }; Ok(response.into()) } diff --git a/crates/contextforge-data-plane-lib/tests/gateway_plugins.rs b/crates/contextforge-data-plane-lib/tests/gateway_plugins.rs index beea51ac..75e8c0b9 100644 --- a/crates/contextforge-data-plane-lib/tests/gateway_plugins.rs +++ b/crates/contextforge-data-plane-lib/tests/gateway_plugins.rs @@ -9,17 +9,20 @@ use cpex::cpex_core::hooks::types::cmf_hook_names; use rmcp::{ ClientHandler, model::{ - CallToolRequestParams, CallToolResult, ClientCapabilities, ClientRequest, ErrorCode, Implementation, - InitializeRequestParams, ProgressNotificationParam, Request, ServerResult, + CallToolRequestParams, CallToolResult, ClientCapabilities, ClientRequest, ContentBlock, ErrorCode, + GetPromptRequestParams, GetPromptResult, Implementation, InitializeRequestParams, ProgressNotificationParam, + Request, ResourceContents, Role as McpRole, ServerResult, }, service::{NotificationContext, PeerRequestOptions, RequestHandle, RoleClient, RunningService}, }; use serde_json::{Map, Value, json}; use support::{ - POST_DENY_ERROR_CODE, PRE_DENY_ERROR_CODE, REWRITTEN_SUM_A, REWRITTEN_SUM_B, RunningGateway, TEST_USER_ID, - TestPlugin, error_code, runtime_with_post, runtime_with_pre, runtime_with_pre_and_post, start_gateway, - start_gateway_with_json_backend_responses, sum_request, text, token, + BACKEND_PROMPT_IMAGE, BACKEND_PROMPT_RESOURCE, POST_DENY_ERROR_CODE, PRE_DENY_ERROR_CODE, PROMPT_ERROR_MESSAGE, + PROMPT_POST_DENY_ERROR_CODE, PromptBehavior, PromptTestPlugin, REWRITTEN_PROMPT_RESOURCE, REWRITTEN_PROMPT_TEXT, + REWRITTEN_PROMPT_TOPIC, REWRITTEN_SUM_A, REWRITTEN_SUM_B, RunningGateway, TEST_USER_ID, TestPlugin, error_code, + error_parts, runtime_with_post, runtime_with_pre, runtime_with_pre_and_post, runtime_with_prompt_plugin, + start_gateway, start_gateway_with_events, start_gateway_with_json_backend_responses, sum_request, text, token, }; type Recorded = Arc>>; @@ -678,3 +681,195 @@ async fn pre_hook_invalid_arguments_return_invalid_params() { assert_eq!(ErrorCode::INVALID_PARAMS, error_code(error)); assert!(gateway.backend_state.calls.lock().expect("backend calls lock poisoned").is_empty()); } + +// --------------------------------------------------------------------------- +// Prompt hooks +// --------------------------------------------------------------------------- + +fn review_request(topic: &str) -> GetPromptRequestParams { + GetPromptRequestParams::new("review") + .with_arguments(serde_json::Map::from_iter([("topic".to_owned(), json!(topic))])) +} + +fn prompt_text(result: &GetPromptResult) -> String { + result + .messages + .iter() + .filter_map(|message| match &message.content { + ContentBlock::Text(text) => Some(text.text.clone()), + _ => None, + }) + .collect() +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 1)] +async fn prompt_pre_hook_rewrites_arguments_reaching_the_backend() { + let plugin = Arc::new(PromptTestPlugin::new("prompt-pre", vec![cmf_hook_names::PROMPT_PRE_FETCH])); + let observations = plugin.observations(); + let runtime = runtime_with_prompt_plugin(plugin).await; + + let gateway = start_gateway(TEST_USER_ID, true, runtime).await; + let service = gateway.connect(TEST_USER_ID).await; + let result = service.get_prompt(review_request("weather")).await.expect("prompt is returned"); + + assert_eq!(format!("review of {REWRITTEN_PROMPT_TOPIC}"), prompt_text(&result)); + + let prompt_calls = gateway.backend_state.prompts.lock().expect("backend prompts lock poisoned"); + assert_eq!("review", prompt_calls[0].tool_name); + assert_eq!( + Some(&Value::from(REWRITTEN_PROMPT_TOPIC)), + prompt_calls[0].args.as_ref().and_then(|args| args.get("topic")) + ); + + let observations = observations.lock().expect("observations lock poisoned"); + assert_eq!(1, observations.pre_calls); + assert_eq!(Some("review"), observations.pre_name.as_deref()); + assert_eq!(Some(gateway.backend_name.as_str()), observations.pre_server_id.as_deref()); +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 1)] +async fn prompt_post_hook_rewrites_rendered_text_before_client_response() { + let plugin = Arc::new(PromptTestPlugin::new("prompt-post", vec![cmf_hook_names::PROMPT_POST_FETCH])); + let observations = plugin.observations(); + let runtime = runtime_with_prompt_plugin(plugin).await; + + let gateway = start_gateway(TEST_USER_ID, true, runtime).await; + let service = gateway.connect(TEST_USER_ID).await; + let result = service.get_prompt(review_request("weather")).await.expect("prompt is returned"); + + assert_eq!(REWRITTEN_PROMPT_TEXT, prompt_text(&result)); + + let prompt_calls = gateway.backend_state.prompts.lock().expect("backend prompts lock poisoned"); + assert_eq!(Some(&Value::from("weather")), prompt_calls[0].args.as_ref().and_then(|args| args.get("topic"))); + drop(prompt_calls); + + let observations = observations.lock().expect("observations lock poisoned"); + assert_eq!(0, observations.pre_calls, "no pre hook is configured"); + assert_eq!(1, observations.post_calls); + assert_eq!(Some("review"), observations.post_prompt_name.as_deref()); +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 1)] +async fn prompt_post_hook_removing_rendered_text_fails_closed() { + let plugin = Arc::new( + PromptTestPlugin::new("prompt-post-drop", vec![cmf_hook_names::PROMPT_POST_FETCH]) + .with_behavior(PromptBehavior::DropText), + ); + let runtime = runtime_with_prompt_plugin(plugin).await; + + let gateway = start_gateway(TEST_USER_ID, true, runtime).await; + let service = gateway.connect(TEST_USER_ID).await; + + let error = service.get_prompt(review_request("weather")).await.expect_err("dropped text fails the call"); + assert_eq!(ErrorCode::INTERNAL_ERROR, error_code(error)); +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 1)] +async fn prompt_post_hook_rewrites_multimodal_prompt_content() { + let plugin = Arc::new(PromptTestPlugin::new("prompt-multimodal", vec![cmf_hook_names::PROMPT_POST_FETCH])); + let runtime = runtime_with_prompt_plugin(plugin).await; + + let gateway = start_gateway(TEST_USER_ID, true, runtime).await; + let service = gateway.connect(TEST_USER_ID).await; + let request = GetPromptRequestParams::new("review_bundle") + .with_arguments(Map::from_iter([("topic".to_owned(), json!("weather"))])); + let result = service.get_prompt(request).await.expect("prompt is returned"); + + assert_eq!(3, result.messages.len()); + assert_eq!(REWRITTEN_PROMPT_TEXT, prompt_text(&result)); + + let ContentBlock::Resource(resource) = &result.messages[1].content else { + panic!("expected the embedded resource to survive as a resource"); + }; + let ResourceContents::TextResourceContents { text, uri, .. } = &resource.resource else { + panic!("expected text resource contents"); + }; + assert_eq!(REWRITTEN_PROMPT_RESOURCE, text, "the plugin's resource edit must reach the client"); + assert_ne!(BACKEND_PROMPT_RESOURCE, text); + assert_eq!("file:///app.env", uri, "identity the plugin did not touch is preserved"); + + let ContentBlock::Image(image) = &result.messages[2].content else { + panic!("expected the image to survive as an image"); + }; + assert_eq!(BACKEND_PROMPT_IMAGE, image.data, "untouched content passes through unchanged"); + assert_eq!(McpRole::Assistant, result.messages[2].role, "roles survive the round trip"); +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 1)] +async fn prompt_post_hook_denial_returns_plugin_error_code() { + let plugin = Arc::new( + PromptTestPlugin::new("prompt-post-deny", vec![cmf_hook_names::PROMPT_POST_FETCH]) + .with_behavior(PromptBehavior::Deny), + ); + let runtime = runtime_with_prompt_plugin(plugin).await; + + let gateway = start_gateway(TEST_USER_ID, true, runtime).await; + let service = gateway.connect(TEST_USER_ID).await; + + let error = service.get_prompt(review_request("weather")).await.expect_err("denied prompt fails the call"); + let (code, message) = error_parts(error); + assert_eq!(ErrorCode(PROMPT_POST_DENY_ERROR_CODE), code); + assert!( + message.contains("prompt"), + "a denied prompt must not be reported to the client as a denied tool call: {message}" + ); +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 1)] +async fn prompt_post_hook_error_flag_fails_the_call() { + let plugin = Arc::new( + PromptTestPlugin::new("prompt-post-error", vec![cmf_hook_names::PROMPT_POST_FETCH]) + .with_behavior(PromptBehavior::MarkError), + ); + let runtime = runtime_with_prompt_plugin(plugin).await; + + let gateway = start_gateway(TEST_USER_ID, true, runtime).await; + let service = gateway.connect(TEST_USER_ID).await; + + let error = service.get_prompt(review_request("weather")).await.expect_err("flagged prompt fails the call"); + let (code, message) = error_parts(error); + assert_eq!(ErrorCode::INVALID_REQUEST, code); + assert_eq!(PROMPT_ERROR_MESSAGE, message); +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 1)] +async fn prompt_hooks_run_either_side_of_the_backend_call() { + let events: Arc>> = Arc::new(StdMutex::new(Vec::new())); + let plugin = Arc::new( + PromptTestPlugin::new( + "prompt-ordering", + vec![cmf_hook_names::PROMPT_PRE_FETCH, cmf_hook_names::PROMPT_POST_FETCH], + ) + .with_events(Arc::clone(&events)), + ); + let runtime = runtime_with_prompt_plugin(plugin).await; + + let gateway = start_gateway_with_events(TEST_USER_ID, runtime, Arc::clone(&events)).await; + let service = gateway.connect(TEST_USER_ID).await; + service.get_prompt(review_request("weather")).await.expect("prompt is returned"); + + assert_eq!(vec!["pre", "backend", "post"], *events.lock().expect("events lock poisoned")); +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 1)] +async fn prompt_pre_and_post_hooks_share_gateway_call_context() { + let plugin = Arc::new( + PromptTestPlugin::new( + "prompt-context", + vec![cmf_hook_names::PROMPT_PRE_FETCH, cmf_hook_names::PROMPT_POST_FETCH], + ) + .with_behavior(PromptBehavior::ContextRoundtrip), + ); + let observations = plugin.observations(); + let runtime = runtime_with_prompt_plugin(plugin).await; + + let gateway = start_gateway(TEST_USER_ID, true, runtime).await; + let service = gateway.connect(TEST_USER_ID).await; + let result = service.get_prompt(review_request("weather")).await.expect("prompt is returned"); + + assert_eq!("review of weather", prompt_text(&result)); + + let observations = observations.lock().expect("observations lock poisoned"); + assert_eq!(1, observations.pre_calls); + assert_eq!(1, observations.post_calls); +} diff --git a/crates/contextforge-data-plane-lib/tests/support/mod.rs b/crates/contextforge-data-plane-lib/tests/support/mod.rs index 54f6898e..a1ce24b3 100644 --- a/crates/contextforge-data-plane-lib/tests/support/mod.rs +++ b/crates/contextforge-data-plane-lib/tests/support/mod.rs @@ -24,9 +24,14 @@ pub(crate) use list_tools_gateway::{ create_tls_gateway_with_four_tls_counters, plaintext_config, }; pub(crate) use plugin::{ - POST_DENY_ERROR_CODE, PRE_DENY_ERROR_CODE, REWRITTEN_SUM_A, REWRITTEN_SUM_B, TestPlugin, TestPluginFactory, + POST_DENY_ERROR_CODE, PRE_DENY_ERROR_CODE, PROMPT_ERROR_MESSAGE, PROMPT_POST_DENY_ERROR_CODE, PromptBehavior, + PromptTestPlugin, REWRITTEN_PROMPT_RESOURCE, REWRITTEN_PROMPT_TEXT, REWRITTEN_PROMPT_TOPIC, REWRITTEN_SUM_A, + REWRITTEN_SUM_B, TestPlugin, TestPluginFactory, }; -pub(crate) use plugin_gateway::{RunningGateway, start_gateway, start_gateway_with_json_backend_responses}; -pub(crate) use runtime::{runtime_with_post, runtime_with_pre, runtime_with_pre_and_post}; -pub(crate) use tool::{error_code, sum_request, text}; +pub(crate) use plugin_gateway::{ + BACKEND_PROMPT_IMAGE, BACKEND_PROMPT_RESOURCE, RunningGateway, start_gateway, start_gateway_with_events, + start_gateway_with_json_backend_responses, +}; +pub(crate) use runtime::{runtime_with_post, runtime_with_pre, runtime_with_pre_and_post, runtime_with_prompt_plugin}; +pub(crate) use tool::{error_code, error_parts, sum_request, text}; pub(crate) use user_config_store::MemoryUserConfigStore; diff --git a/crates/contextforge-data-plane-lib/tests/support/plugin.rs b/crates/contextforge-data-plane-lib/tests/support/plugin.rs index ddc02109..788bfe9f 100644 --- a/crates/contextforge-data-plane-lib/tests/support/plugin.rs +++ b/crates/contextforge-data-plane-lib/tests/support/plugin.rs @@ -5,7 +5,7 @@ use std::{ use async_trait::async_trait; use cpex::cpex_core::{ - cmf::{CmfHook, ContentPart, Message, MessagePayload, Role}, + cmf::{CmfHook, ContentPart, Message, MessagePayload, PromptResult as CmfPromptResult, Role}, context::PluginContext, error::{PluginError, PluginViolation}, factory::{PluginFactory, PluginInstance}, @@ -314,6 +314,198 @@ impl TestPluginFactory { } } +pub(crate) const REWRITTEN_PROMPT_TOPIC: &str = "rewritten-topic"; +pub(crate) const REWRITTEN_PROMPT_TEXT: &str = "review of [REDACTED]"; +pub(crate) const REWRITTEN_PROMPT_RESOURCE: &str = "config with [REDACTED]"; +pub(crate) const PROMPT_POST_DENY_ERROR_CODE: i32 = -32004; +pub(crate) const PROMPT_ERROR_MESSAGE: &str = "prompt blocked by policy"; + +fn prompt_result_mut(payload: &mut MessagePayload) -> Option<&mut CmfPromptResult> { + payload.message.content.iter_mut().find_map(|part| match part { + ContentPart::PromptResult { content } => Some(content), + _ => None, + }) +} + +#[derive(Clone, Copy, Default)] +pub(crate) enum PromptBehavior { + #[default] + Rewrite, + DropText, + ContextRoundtrip, + Deny, + MarkError, +} + +pub(crate) struct PromptTestPlugin { + pub(crate) config: PluginConfig, + pub(crate) observations: Arc>, + pub(crate) behavior: PromptBehavior, + pub(crate) events: Option>>>, +} + +#[derive(Default)] +pub(crate) struct PromptObservations { + pub(crate) pre_calls: usize, + pub(crate) pre_name: Option, + pub(crate) pre_server_id: Option, + pub(crate) post_calls: usize, + pub(crate) post_prompt_name: Option, +} + +impl PromptTestPlugin { + pub(crate) fn new(name: &str, hooks: Vec<&'static str>) -> Self { + Self { + config: PluginConfig { + name: name.to_owned(), + kind: "prompt-test".to_owned(), + hooks: hooks.into_iter().map(str::to_owned).collect(), + ..Default::default() + }, + observations: Arc::new(Mutex::new(PromptObservations::default())), + behavior: PromptBehavior::default(), + events: None, + } + } + + pub(crate) fn with_events(mut self, events: Arc>>) -> Self { + self.events = Some(events); + self + } + + pub(crate) fn rebuild(&self, config: PluginConfig) -> Self { + Self { + config, + observations: Arc::clone(&self.observations), + behavior: self.behavior, + events: self.events.clone(), + } + } + + fn record(&self, event: &'static str) { + if let Some(events) = &self.events { + events.lock().expect("events lock poisoned").push(event); + } + } + + pub(crate) fn with_behavior(mut self, behavior: PromptBehavior) -> Self { + self.behavior = behavior; + self + } + + pub(crate) fn observations(&self) -> Arc> { + Arc::clone(&self.observations) + } + + fn handle_pre(&self, payload: &MessagePayload, ctx: &mut PluginContext) -> PluginResult { + self.record("pre"); + let mut observations = self.observations.lock().expect("observations lock poisoned"); + observations.pre_calls += 1; + if let Some(request) = payload.message.get_prompt_requests().first() { + observations.pre_name = Some(request.name.clone()); + observations.pre_server_id.clone_from(&request.server_id); + } + drop(observations); + + match self.behavior { + PromptBehavior::ContextRoundtrip => { + ctx.set_global("prompt_pre_seen", json!(true)); + PluginResult::allow() + }, + PromptBehavior::Rewrite | PromptBehavior::DropText | PromptBehavior::Deny | PromptBehavior::MarkError => { + let mut modified = payload.clone(); + if let Some(ContentPart::PromptRequest { content }) = + modified.message.content.iter_mut().find(|part| matches!(part, ContentPart::PromptRequest { .. })) + { + content.arguments = HashMap::from([("topic".to_owned(), json!(REWRITTEN_PROMPT_TOPIC))]); + } + PluginResult::modify_payload(modified) + }, + } + } + + fn handle_post(&self, payload: &MessagePayload, ctx: &mut PluginContext) -> PluginResult { + self.record("post"); + let mut observations = self.observations.lock().expect("observations lock poisoned"); + observations.post_calls += 1; + if let Some(result) = payload.message.get_prompt_results().first() { + observations.post_prompt_name = Some(result.prompt_name.clone()); + } + drop(observations); + + match self.behavior { + PromptBehavior::Rewrite => { + let mut modified = payload.clone(); + if let Some(result) = prompt_result_mut(&mut modified) { + for message in &mut result.messages { + for part in &mut message.content { + match part { + ContentPart::Text { text } => REWRITTEN_PROMPT_TEXT.clone_into(text), + ContentPart::Resource { content } => { + content.content = Some(REWRITTEN_PROMPT_RESOURCE.to_owned()); + }, + _ => {}, + } + } + } + } + PluginResult::modify_payload(modified) + }, + PromptBehavior::DropText => { + let mut modified = payload.clone(); + if let Some(result) = prompt_result_mut(&mut modified) { + result.messages.clear(); + } + PluginResult::modify_payload(modified) + }, + PromptBehavior::MarkError => { + let mut modified = payload.clone(); + if let Some(result) = prompt_result_mut(&mut modified) { + result.is_error = true; + result.error_message = Some(PROMPT_ERROR_MESSAGE.to_owned()); + } + PluginResult::modify_payload(modified) + }, + PromptBehavior::Deny => PluginResult::deny( + PluginViolation::new("prompt_post_denied", "prompt post denied") + .with_proto_error_code(i64::from(PROMPT_POST_DENY_ERROR_CODE)), + ), + PromptBehavior::ContextRoundtrip => { + if ctx.get_global("prompt_pre_seen") == Some(&json!(true)) { + PluginResult::allow() + } else { + PluginResult::deny( + PluginViolation::new("missing_prompt_context", "prompt pre context missing") + .with_proto_error_code(i64::from(MISSING_CONTEXT_ERROR_CODE)), + ) + } + }, + } + } +} + +#[async_trait] +impl Plugin for PromptTestPlugin { + fn config(&self) -> &PluginConfig { + &self.config + } +} + +impl HookHandler for PromptTestPlugin { + async fn handle( + &self, + payload: &MessagePayload, + _extensions: &Extensions, + ctx: &mut PluginContext, + ) -> PluginResult { + if payload.message.get_prompt_results().is_empty() { + self.handle_pre(payload, ctx) + } else { + self.handle_post(payload, ctx) + } + } +} + impl PluginFactory for TestPluginFactory { fn create(&self, config: &PluginConfig) -> Result> { let plugin = Arc::new(TestPlugin { diff --git a/crates/contextforge-data-plane-lib/tests/support/plugin_gateway.rs b/crates/contextforge-data-plane-lib/tests/support/plugin_gateway.rs index 381229cb..a118e25a 100644 --- a/crates/contextforge-data-plane-lib/tests/support/plugin_gateway.rs +++ b/crates/contextforge-data-plane-lib/tests/support/plugin_gateway.rs @@ -15,9 +15,9 @@ use http::{HeaderMap, HeaderValue}; use rmcp::{ ErrorData, RoleClient, RoleServer, ServerHandler, ServiceExt, model::{ - CallToolRequestParams, CallToolResponse, CallToolResult, ContentBlock, ErrorCode, Implementation, - InitializeRequestParams, InitializeResult, NumberOrString, ProgressNotificationParam, ProgressToken, - ServerCapabilities, + CallToolRequestParams, CallToolResponse, CallToolResult, ContentBlock, ErrorCode, GetPromptRequestParams, + GetPromptResponse, GetPromptResult, Implementation, InitializeRequestParams, InitializeResult, NumberOrString, + ProgressNotificationParam, ProgressToken, PromptMessage, ResourceContents, Role, ServerCapabilities, }, service::{RequestContext, Service}, transport::{ @@ -31,6 +31,9 @@ use tokio::sync::Mutex as TokioMutex; use super::{MemoryUserConfigStore, token}; +pub(crate) const BACKEND_PROMPT_RESOURCE: &str = "token=secret"; +pub(crate) const BACKEND_PROMPT_IMAGE: &str = "aW1hZ2UtYnl0ZXM="; + static GATEWAY_PORT_LOCK: OnceLock>> = OnceLock::new(); const CLIENT_CONNECT_TIMEOUT: Duration = Duration::from_secs(2); const GATEWAY_PORT_READY_TIMEOUT: Duration = Duration::from_secs(10); @@ -45,7 +48,9 @@ pub(crate) struct BackendObservation { #[derive(Clone, Default)] pub(crate) struct BackendState { pub(crate) calls: Arc>>, + pub(crate) prompts: Arc>>, pub(crate) cancellations: Arc>>, + pub(crate) events: Arc>>, } #[derive(Clone)] @@ -59,10 +64,43 @@ impl ServerHandler for TestBackend { _request: InitializeRequestParams, _cx: RequestContext, ) -> Result { - Ok(InitializeResult::new(ServerCapabilities::builder().enable_tools().build()) + Ok(InitializeResult::new(ServerCapabilities::builder().enable_tools().enable_prompts().build()) .with_server_info(Implementation::new("test-backend", "0.1.0"))) } + async fn get_prompt( + &self, + request: GetPromptRequestParams, + _cx: RequestContext, + ) -> Result { + self.state + .prompts + .lock() + .expect("backend prompts lock poisoned") + .push(BackendObservation { tool_name: request.name.clone(), args: request.arguments.clone() }); + self.state.events.lock().expect("backend events lock poisoned").push("backend"); + + let topic = request + .arguments + .as_ref() + .and_then(|arguments| arguments.get("topic")) + .and_then(Value::as_str) + .unwrap_or("nothing"); + if request.name == "review_bundle" { + return Ok(GetPromptResult::new(vec![ + PromptMessage::new_text(Role::User, format!("review of {topic}")), + PromptMessage::new( + Role::User, + ContentBlock::resource(ResourceContents::text(BACKEND_PROMPT_RESOURCE, "file:///app.env")), + ), + PromptMessage::new(Role::Assistant, ContentBlock::image(BACKEND_PROMPT_IMAGE, "image/png")), + ]) + .into()); + } + + Ok(GetPromptResult::new(vec![PromptMessage::new_text(Role::User, format!("review of {topic}"))]).into()) + } + async fn call_tool( &self, request: CallToolRequestParams, @@ -222,6 +260,15 @@ pub(crate) async fn start_gateway( start_gateway_with_runtime(user, runtime_plugins_enabled, plugin_runtime, false).await } +pub(crate) async fn start_gateway_with_events( + user: &str, + plugin_runtime: Arc, + events: Arc>>, +) -> RunningGateway { + start_gateway_with_state(user, true, plugin_runtime, false, BackendState { events, ..BackendState::default() }) + .await +} + pub(crate) async fn start_gateway_with_json_backend_responses( user: &str, runtime_plugins_enabled: bool, @@ -235,6 +282,23 @@ async fn start_gateway_with_runtime( runtime_plugins_enabled: bool, plugin_runtime: Arc, json_backend_responses: bool, +) -> RunningGateway { + start_gateway_with_state( + user, + runtime_plugins_enabled, + plugin_runtime, + json_backend_responses, + BackendState::default(), + ) + .await +} + +async fn start_gateway_with_state( + user: &str, + runtime_plugins_enabled: bool, + plugin_runtime: Arc, + json_backend_responses: bool, + backend_state: BackendState, ) -> RunningGateway { let port_lock = Arc::clone(GATEWAY_PORT_LOCK.get_or_init(|| Arc::new(TokioMutex::new(())))); let port_guard = port_lock.lock().await; @@ -243,7 +307,6 @@ async fn start_gateway_with_runtime( let backend_port = backend_listener.local_addr().expect("backend address").port(); let backend_name = format!("backend-{backend_port}"); let virtual_host_id = "vh-cpex-test"; - let backend_state = BackendState::default(); let backend_service = StreamableHttpService::new( { diff --git a/crates/contextforge-data-plane-lib/tests/support/runtime.rs b/crates/contextforge-data-plane-lib/tests/support/runtime.rs index 83d9b3dd..fc4c6c0b 100644 --- a/crates/contextforge-data-plane-lib/tests/support/runtime.rs +++ b/crates/contextforge-data-plane-lib/tests/support/runtime.rs @@ -4,7 +4,27 @@ use contextforge_data_plane_cpex::CpexRuntimeRegistry; use cpex::cpex_core::config::CpexConfig; use serde_json::json; -use super::{TestPlugin, TestPluginFactory}; +use contextforge_data_plane_cpex::CmfPluginFactory; + +use super::{PromptTestPlugin, TestPlugin, TestPluginFactory}; + +pub(crate) async fn runtime_with_prompt_plugin(plugin: Arc) -> Arc { + let mut runtime = CpexRuntimeRegistry::default(); + let template = Arc::clone(&plugin); + runtime + .register_factory("prompt-test", Box::new(CmfPluginFactory::new(move |config| template.rebuild(config)))) + .expect("prompt test factory registers"); + let config = serde_json::from_value(json!({ + "plugins": [{ + "name": plugin.config.name.clone(), + "kind": plugin.config.kind.clone(), + "hooks": plugin.config.hooks.clone(), + }] + })) + .expect("prompt CPEX config parses"); + runtime.apply_config(Some(config)).await.expect("prompt runtime config applies"); + Arc::new(runtime) +} pub(crate) async fn runtime_with_pre(plugin: Arc) -> Arc { runtime_with_plugins(&[plugin]).await diff --git a/crates/contextforge-data-plane-lib/tests/support/tool.rs b/crates/contextforge-data-plane-lib/tests/support/tool.rs index cc36c55b..e4793575 100644 --- a/crates/contextforge-data-plane-lib/tests/support/tool.rs +++ b/crates/contextforge-data-plane-lib/tests/support/tool.rs @@ -26,3 +26,10 @@ pub(crate) fn error_code(error: ServiceError) -> ErrorCode { }; error.code } + +pub(crate) fn error_parts(error: ServiceError) -> (ErrorCode, String) { + let ServiceError::McpError(error) = error else { + panic!("expected MCP error, got {error:?}"); + }; + (error.code, error.message.into_owned()) +}