From c2b4d4d2c93d659046e98befee9c7f329966544f Mon Sep 17 00:00:00 2001 From: lucarlig Date: Mon, 7 Sep 2026 10:53:56 +0100 Subject: [PATCH] refactor: simplify CPEX hook execution and request state Signed-off-by: lucarlig --- _context/wiki/routing.md | 29 +- .../contextforge-data-plane-cpex/src/cmf.rs | 1071 +-------------- .../src/factory.rs | 12 +- .../src/handle.rs | 1191 ----------------- .../contextforge-data-plane-cpex/src/hooks.rs | 59 +- .../contextforge-data-plane-cpex/src/lib.rs | 16 +- .../src/pipeline.rs | 162 --- .../src/prompts/mod.rs | 287 ++++ .../src/prompts/tests.rs | 436 ++++++ .../src/registry/mod.rs | 181 +++ .../src/registry/tests.rs | 976 ++++++++++++++ .../src/resources/mod.rs | 161 +++ .../src/resources/tests.rs | 142 ++ .../src/runtime.rs | 381 ++---- .../src/tools/mod.rs | 187 +++ .../src/tools/tests.rs | 16 + .../src/gateway/backend_client.rs | 24 +- .../src/gateway/mcp_service/initialization.rs | 2 +- .../src/gateway/mcp_service/prompts.rs | 14 +- .../src/gateway/mcp_service/tools.rs | 24 +- 20 files changed, 2612 insertions(+), 2759 deletions(-) delete mode 100644 crates/contextforge-data-plane-cpex/src/handle.rs delete mode 100644 crates/contextforge-data-plane-cpex/src/pipeline.rs create mode 100644 crates/contextforge-data-plane-cpex/src/prompts/mod.rs create mode 100644 crates/contextforge-data-plane-cpex/src/prompts/tests.rs create mode 100644 crates/contextforge-data-plane-cpex/src/registry/mod.rs create mode 100644 crates/contextforge-data-plane-cpex/src/registry/tests.rs create mode 100644 crates/contextforge-data-plane-cpex/src/resources/mod.rs create mode 100644 crates/contextforge-data-plane-cpex/src/resources/tests.rs create mode 100644 crates/contextforge-data-plane-cpex/src/tools/mod.rs create mode 100644 crates/contextforge-data-plane-cpex/src/tools/tests.rs diff --git a/_context/wiki/routing.md b/_context/wiki/routing.md index f62579fd..0dc2e105 100644 --- a/_context/wiki/routing.md +++ b/_context/wiki/routing.md @@ -49,4 +49,31 @@ For clients on `≥ 2026-07-28`, `call_tool` validates `Mcp-Param-*` headers aga ## Plugin hooks -`call_tool` and `get_prompt` run `before_*/after_*` hooks when a `GatewayPluginRuntimeHandle` is configured. Pre-hook may rewrite arguments or deny; post-hook may rewrite or reject the response. Pre-hook state is passed to the post-hook. +`call_tool`, `get_prompt`, and `read_resource` run pre/post hooks when a +`GatewayPluginRuntimeHandle` is configured. Pre-hooks can deny a request or edit +tool/prompt arguments or the resource URI. Resource URI edits must resolve through +the caller's published routes. Post-hooks can rewrite or reject the response. + +The handle selects a runtime before backend I/O. It returns typed request state +whose `after_*` method runs the post-hook on that same runtime, even after a reload +or a reload failure. A request that started without post-hooks never gains one +mid-flight. Tool state is shared under a mutex so progress notifications and the +final response use the same correlation ID and serialize plugin-context updates. +Prompt and resource state is owned by a single request and needs no mutex. + +The internal CPEX crate separates these responsibilities: + +| Module | Owns | +| --- | --- | +| `registry/` | Runtime selection, factory registration, config reloads, and the watcher | +| `runtime.rs` | Manager lifecycle, shared pre/post execution, and request context | +| `tools/`, `prompts/`, `resources/` | Typed request state, operation-specific CMF conversion, and conversion tests | +| `cmf.rs` | Hook inventory, message envelopes, and the `CmfResponse` conversion trait | +| `hooks.rs` | Shared argument edits and pre-hook results | +| `config.rs`, `factory.rs` | Config decoding/storage and compiled-in plugin factories | + +Each response adapter implements `CmfResponse`; the runner handles invocation, +unchanged payloads, context propagation, and denial. Conversion and rejection +rules remain operation-specific: tools, rendered prompts, and resource reads +have different MCP representations. See [Config](config.md#tool-call-hook-behavior) +for those contracts. diff --git a/crates/contextforge-data-plane-cpex/src/cmf.rs b/crates/contextforge-data-plane-cpex/src/cmf.rs index 010606c8..0f24a6a7 100644 --- a/crates/contextforge-data-plane-cpex/src/cmf.rs +++ b/crates/contextforge-data-plane-cpex/src/cmf.rs @@ -1,1044 +1,73 @@ -use std::collections::HashMap; +//! Common CMF envelopes and the response contract used by the hook runner. +//! Conversion rules stay with each operation because their MCP semantics differ. -use base64::{Engine as _, prelude::BASE64_STANDARD}; -use cpex::cpex_core::cmf::{ - AudioSource, ContentPart, ImageSource, Message, MessagePayload, PromptRequest, PromptResult, - Resource as CmfResource, ResourceReference, ResourceType, Role, ToolCall, ToolResult, constants::SCHEMA_VERSION, +use cpex::cpex_core::{ + cmf::{ContentPart, Message, MessagePayload, Role, constants::SCHEMA_VERSION}, + executor::PipelineResult, + hooks::types::cmf_hook_names, }; -use rmcp::model::{ - CallToolRequestParams, CallToolResult, ContentBlock, GetPromptRequestParams, GetPromptResult, PromptMessage, - ReadResourceResult, Resource as McpResource, ResourceContents, Role as McpRole, -}; -use serde_json::{Map, Value}; - -pub(crate) fn tool_call_payload( - request: &CallToolRequestParams, - tool_name: &str, - backend_name: &str, - tool_call_id: &str, -) -> MessagePayload { - MessagePayload { - message: Message { - schema_version: SCHEMA_VERSION.to_owned(), - role: Role::Assistant, - content: vec![ContentPart::ToolCall { - content: ToolCall { - tool_call_id: tool_call_id.to_owned(), - name: tool_name.to_owned(), - arguments: request.arguments.clone().unwrap_or_default().into_iter().collect(), - namespace: Some(backend_name.to_owned()), - }, - }], - channel: None, - }, - } -} - -pub(crate) fn resource_request_payload(resource_uri: &str, resource_request_id: &str) -> MessagePayload { - MessagePayload { - message: Message { - schema_version: SCHEMA_VERSION.to_owned(), - role: Role::User, - content: vec![ContentPart::ResourceRef { - content: ResourceReference { - resource_request_id: resource_request_id.to_owned(), - uri: resource_uri.to_owned(), - name: None, - resource_type: ResourceType::Uri, - range_start: None, - range_end: None, - selector: None, - }, - }], - channel: None, - }, - } -} - -pub(crate) fn resource_result_payload( - response: &ReadResourceResult, - resource_request_id: &str, -) -> Option { - let content = response - .contents - .iter() - .map(|content| { - cmf_resource_content(content, resource_request_id).map(|content| ContentPart::Resource { content }) - }) - .collect::>>()?; - Some(MessagePayload { - message: Message { schema_version: SCHEMA_VERSION.to_owned(), role: Role::Assistant, content, channel: None }, - }) -} - -fn cmf_resource_content(content: &ResourceContents, resource_request_id: &str) -> Option { - let (uri, mime_type, text, blob) = match content { - ResourceContents::TextResourceContents { uri, mime_type, text, .. } => { - (uri.clone(), mime_type.clone(), Some(text.clone()), None) - }, - ResourceContents::BlobResourceContents { uri, mime_type, blob, .. } => { - (uri.clone(), mime_type.clone(), None, Some(BASE64_STANDARD.decode(blob).ok()?)) - }, - _ => return None, - }; - Some(CmfResource { - resource_request_id: resource_request_id.to_owned(), - uri, - resource_type: ResourceType::Uri, - content: text, - blob, - mime_type, - ..Default::default() - }) -} - -pub(crate) fn resource_result_response( - mut original: ReadResourceResult, - payload: &MessagePayload, -) -> Option { - // Resource post hooks replace each resource's content, not the read envelope. - if payload.message.content.len() != original.contents.len() { - return None; - } - for (original, modified) in original.contents.iter_mut().zip(&payload.message.content) { - let ContentPart::Resource { content } = modified else { return None }; - let meta = match original { - ResourceContents::TextResourceContents { meta, .. } - | ResourceContents::BlobResourceContents { meta, .. } => meta.clone(), - _ => return None, - }; - *original = match (&content.content, &content.blob) { - (Some(text), _) => ResourceContents::TextResourceContents { - uri: content.uri.clone(), - mime_type: content.mime_type.clone(), - text: text.clone(), - meta, - }, - (None, Some(bytes)) => { - let blob = match original { - ResourceContents::BlobResourceContents { blob, .. } - if BASE64_STANDARD.decode(blob.as_bytes()).ok().as_ref() == Some(bytes) => - { - blob.clone() - }, - _ => BASE64_STANDARD.encode(bytes), - }; - ResourceContents::BlobResourceContents { - uri: content.uri.clone(), - mime_type: content.mime_type.clone(), - blob, - meta, - } - }, - _ => return None, - }; - } - Some(original) -} - -pub(crate) fn tool_result_payload(tool_name: &str, response: &CallToolResult, tool_call_id: &str) -> MessagePayload { - tool_json_result_payload( - tool_name, - serde_json::to_value(response).unwrap_or(Value::Null), - response.is_error.unwrap_or(false), - tool_call_id, - ) -} - -pub(crate) fn tool_json_result_payload( - tool_name: &str, - content: Value, - is_error: bool, - tool_call_id: &str, -) -> MessagePayload { - MessagePayload { - message: Message { - schema_version: SCHEMA_VERSION.to_owned(), - role: Role::Tool, - content: vec![ContentPart::ToolResult { - content: ToolResult { - tool_call_id: tool_call_id.to_owned(), - tool_name: tool_name.to_owned(), - content, - is_error, - }, - }], - channel: None, - }, - } -} - -pub(crate) fn tool_result_content(payload: &MessagePayload) -> Option { - payload.message.get_tool_results().first().map(|tool_result| tool_result.content.clone()) -} - -pub(crate) fn tool_call_arguments(payload: &MessagePayload) -> Option> { - payload - .message - .get_tool_calls() - .first() - .map(|tool_call| tool_call.arguments.clone().into_iter().collect::>()) -} - -pub(crate) fn tool_result_response(original: CallToolResult, payload: &MessagePayload) -> CallToolResult { - let mut result = payload.message.get_tool_results().first().map_or(original, |tool_result| { - serde_json::from_value::(tool_result.content.clone()).map_or_else( - |_| { - if tool_result.is_error { - raw_error_tool_result(tool_result.content.clone()) - } else { - raw_success_tool_result(tool_result.content.clone()) - } - }, - |mut result| { - result.is_error = Some(tool_result.is_error); - result - }, - ) - }); +use rmcp::{ErrorData, model::ErrorCode}; - let text = payload.message.get_text_content(); - if !text.is_empty() { - result.content.push(ContentBlock::text(text)); - } - - result -} - -fn raw_success_tool_result(value: Value) -> CallToolResult { - if let Value::String(text) = value { - CallToolResult::success(vec![ContentBlock::text(text)]) - } else { - CallToolResult::structured(value) - } -} - -fn raw_error_tool_result(value: Value) -> CallToolResult { - if let Value::String(text) = value { - CallToolResult::error(vec![ContentBlock::text(text)]) - } else { - CallToolResult::structured_error(value) - } -} - -pub(crate) fn prompt_request_payload( - request: &GetPromptRequestParams, - prompt_name: &str, - backend_name: &str, - prompt_request_id: &str, -) -> MessagePayload { - MessagePayload { - message: Message { - schema_version: SCHEMA_VERSION.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::>()) +#[derive(Clone, Copy)] +pub(crate) enum Operation { + Tool, + Prompt, + Resource, } -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, - }, - } -} +impl Operation { + pub(crate) const ALL: [Self; 3] = [Self::Tool, Self::Prompt, Self::Resource]; -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; + pub(crate) fn hooks(self) -> [&'static str; 2] { + match self { + Self::Tool => [cmf_hook_names::TOOL_PRE_INVOKE, cmf_hook_names::TOOL_POST_INVOKE], + Self::Prompt => [cmf_hook_names::PROMPT_PRE_FETCH, cmf_hook_names::PROMPT_POST_FETCH], + Self::Resource => [cmf_hook_names::RESOURCE_PRE_FETCH, cmf_hook_names::RESOURCE_POST_FETCH], } + } - 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; + pub(crate) fn subject(self) -> &'static str { + match self { + Self::Tool => "tool call", + Self::Prompt => "prompt", + Self::Resource => "resource", } - *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, + pub(crate) fn id_prefix(self) -> &'static str { + match self { + Self::Tool => "gateway-tool-call", + Self::Prompt => "gateway-prompt-request", + Self::Resource => "gateway-resource-request", + } } } -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, - }; +/// Only the projection and application differ between response types. +/// The runner owns invocation, unchanged payloads, context, and denial handling. +pub(crate) trait CmfResponse: Sized { + const OPERATION: Operation; - Some(part) + fn to_payload(&self, name: &str, id: &str) -> Result; + fn apply_payload(self, payload: &MessagePayload, name: &str, id: &str) -> Result; } -// 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) +pub(crate) fn message_payload(role: Role, content: Vec) -> MessagePayload { + MessagePayload { message: Message { schema_version: SCHEMA_VERSION.to_owned(), role, content, channel: None } } } -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)) +pub(crate) fn modified_message_payload(result: &PipelineResult) -> Option<&MessagePayload> { + result.modified_payload.as_ref().and_then(|payload| payload.as_any().downcast_ref::()) } -#[cfg(test)] -mod tests { - use super::*; - - #[test] - fn resource_result_response_applies_text_changes() { - let original = - ReadResourceResult::new(vec![ResourceContents::text("AWS_ACCESS_KEY_ID=secret", "file:///password.env")]); - let mut payload = resource_result_payload(&original, "resource-1").expect("resource result is supported"); - let ContentPart::Resource { content } = &mut payload.message.content[0] else { - panic!("expected resource content"); - }; - content.content = Some("AWS_ACCESS_KEY_ID=[redacted]".to_owned()); - - let result = resource_result_response(original, &payload).expect("text edit applies"); - - let ResourceContents::TextResourceContents { text, uri, .. } = &result.contents[0] else { - panic!("expected text resource"); - }; - assert_eq!("AWS_ACCESS_KEY_ID=[redacted]", text); - assert_eq!("file:///password.env", uri); - } - - #[test] - fn resource_result_response_decodes_and_applies_blob_changes() { - let wire_blob = BASE64_STANDARD.encode(b"AWS_ACCESS_KEY_ID=secret"); - let original = ReadResourceResult::new(vec![ - ResourceContents::blob(wire_blob.clone(), "file:///password.bin") - .with_mime_type("application/octet-stream"), - ]); - let mut payload = resource_result_payload(&original, "resource-1").expect("valid blob is supported"); - let ContentPart::Resource { content } = &mut payload.message.content[0] else { - panic!("expected resource content"); - }; - assert_eq!(Some(b"AWS_ACCESS_KEY_ID=secret".as_slice()), content.blob.as_deref()); - content.blob = Some(b"AWS_ACCESS_KEY_ID=[redacted]".to_vec()); - - let result = resource_result_response(original, &payload).expect("blob edit applies"); - - let ResourceContents::BlobResourceContents { blob, uri, .. } = &result.contents[0] else { - panic!("expected blob resource"); - }; - assert_eq!(b"AWS_ACCESS_KEY_ID=[redacted]", BASE64_STANDARD.decode(blob).expect("valid base64").as_slice()); - assert_eq!("file:///password.bin", uri); - assert_ne!(&wire_blob, blob); - } - - #[test] - fn resource_result_response_preserves_unchanged_blob_wire_value() { - let wire_blob = BASE64_STANDARD.encode(b"unchanged"); - let original = ReadResourceResult::new(vec![ResourceContents::blob(&wire_blob, "file:///image.bin")]); - let payload = resource_result_payload(&original, "resource-1").expect("valid blob is supported"); - - let result = resource_result_response(original, &payload).expect("unchanged blob applies"); - - let ResourceContents::BlobResourceContents { blob, .. } = &result.contents[0] else { - panic!("expected blob resource"); - }; - assert_eq!(&wire_blob, blob); - } - - #[test] - fn resource_result_payload_rejects_invalid_base64_blob() { - let original = ReadResourceResult::new(vec![ResourceContents::blob("not base64!", "file:///image.bin")]); - - assert!(resource_result_payload(&original, "resource-1").is_none()); - } - - #[test] - fn resource_result_allows_mime_uri_and_content_type_changes() { - let original: ReadResourceResult = serde_json::from_value(serde_json::json!({ - "_meta": {"response": "preserved"}, - "contents": [ - {"uri": "file:///a", "mimeType": "text/plain", "text": "original", "_meta": {"item": 1}}, - {"uri": "file:///b", "mimeType": "application/octet-stream", "blob": "YmluYXJ5", "_meta": {"item": 2}} - ] - })) - .expect("valid resource response"); - let mut payload = resource_result_payload(&original, "resource-1").expect("resource payload"); - for (index, part) in payload.message.content.iter_mut().enumerate() { - let ContentPart::Resource { content } = part else { panic!("resource content") }; - content.uri = format!("file:///changed-{index}"); - if index == 0 { - content.content = None; - content.blob = Some(b"binary edit".to_vec()); - content.mime_type = Some("application/octet-stream".to_owned()); - } else { - content.blob = None; - content.content = Some("text edit".to_owned()); - content.mime_type = Some("text/plain".to_owned()); - } - } - let actual = serde_json::to_value(resource_result_response(original, &payload).expect("valid changes apply")) - .expect("response serializes"); - assert_eq!(serde_json::json!({"response": "preserved"}), actual["_meta"]); - assert_eq!("file:///changed-0", actual["contents"][0]["uri"]); - assert_eq!("application/octet-stream", actual["contents"][0]["mimeType"]); - assert_eq!(BASE64_STANDARD.encode(b"binary edit"), actual["contents"][0]["blob"]); - assert_eq!(1, actual["contents"][0]["_meta"]["item"]); - assert_eq!("file:///changed-1", actual["contents"][1]["uri"]); - assert_eq!("text/plain", actual["contents"][1]["mimeType"]); - assert_eq!("text edit", actual["contents"][1]["text"]); - assert_eq!(2, actual["contents"][1]["_meta"]["item"]); - } - - #[test] - fn resource_result_ignores_cmf_fields_that_are_not_mcp_content() { - let original = ReadResourceResult::new(vec![ResourceContents::text("original", "file:///a")]); - let mut payload = resource_result_payload(&original, "resource-1").expect("resource payload"); - payload.message.schema_version = "plugin value".to_owned(); - payload.message.role = Role::User; - payload.message.channel = Some(cpex::cpex_core::cmf::Channel::Analysis); - let ContentPart::Resource { content } = &mut payload.message.content[0] else { panic!("resource content") }; - content.resource_request_id = "plugin value".to_owned(); - content.name = Some("display name".to_owned()); - content.description = Some("description".to_owned()); - content.size_bytes = Some(8); - content.version = Some("v2".to_owned()); - content.annotations.insert("note".to_owned(), serde_json::json!("annotation")); - content.content = Some("redacted".to_owned()); - let result = resource_result_response(original, &payload).expect("MCP content remains usable"); - let ResourceContents::TextResourceContents { text, .. } = &result.contents[0] else { panic!("text content") }; - assert_eq!("redacted", text); - } - - #[test] - fn resource_result_prefers_text_when_both_content_fields_are_present() { - let original = ReadResourceResult::new(vec![ResourceContents::blob("YmluYXJ5", "file:///a")]); - let mut payload = resource_result_payload(&original, "resource-1").expect("resource payload"); - let ContentPart::Resource { content } = &mut payload.message.content[0] else { panic!("resource content") }; - content.content = Some("text replacement".to_owned()); - let result = - resource_result_response(original, &payload).expect("text takes precedence, as in the built-in serializer"); - let ResourceContents::TextResourceContents { text, .. } = &result.contents[0] else { panic!("text resource") }; - assert_eq!("text replacement", text); - } - - #[test] - fn resource_result_rejects_content_without_a_valid_mcp_representation() { - let original = ReadResourceResult::new(vec![ResourceContents::text("original", "file:///a")]); - let mut payload = resource_result_payload(&original, "resource-1").expect("resource payload"); - let ContentPart::Resource { content } = &mut payload.message.content[0] else { panic!("resource content") }; - content.content = None; - assert!(resource_result_response(original, &payload).is_none()); - } - - 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")]); - let nested = CallToolResult::success(vec![ContentBlock::text("changed")]); - let mut payload = tool_result_payload("sum", &nested, "call-1"); - let ContentPart::ToolResult { content } = &mut payload.message.content[0] else { - panic!("expected tool result"); - }; - content.is_error = true; - - let result = tool_result_response(original, &payload); +pub(crate) fn plugin_denied_error(subject: &str, result: PipelineResult) -> ErrorData { + let code = result + .violation + .and_then(|violation| { + tracing::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); - assert_eq!(Some(true), result.is_error); - } + ErrorData { code, message: format!("Plugin denied {subject}").into(), data: None } } diff --git a/crates/contextforge-data-plane-cpex/src/factory.rs b/crates/contextforge-data-plane-cpex/src/factory.rs index 627abe1c..3034fbbd 100644 --- a/crates/contextforge-data-plane-cpex/src/factory.rs +++ b/crates/contextforge-data-plane-cpex/src/factory.rs @@ -4,7 +4,7 @@ use cpex::cpex_core::{ cmf::CmfHook, error::PluginError, factory::{PluginFactory, PluginInstance}, - hooks::{HookHandler, TypedHandlerAdapter, types::cmf_hook_names}, + hooks::{HookHandler, TypedHandlerAdapter}, plugin::{Plugin, PluginConfig}, registry::AnyHookHandler, }; @@ -46,13 +46,5 @@ where } pub(crate) fn supported_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), - cmf_hook_names::RESOURCE_PRE_FETCH => Some(cmf_hook_names::RESOURCE_PRE_FETCH), - cmf_hook_names::RESOURCE_POST_FETCH => Some(cmf_hook_names::RESOURCE_POST_FETCH), - _ => None, - } + crate::cmf::Operation::ALL.into_iter().flat_map(crate::cmf::Operation::hooks).find(|name| *name == hook) } diff --git a/crates/contextforge-data-plane-cpex/src/handle.rs b/crates/contextforge-data-plane-cpex/src/handle.rs deleted file mode 100644 index f6e8e507..00000000 --- a/crates/contextforge-data-plane-cpex/src/handle.rs +++ /dev/null @@ -1,1191 +0,0 @@ -use std::{ - sync::{ - Arc, - atomic::{AtomicBool, Ordering}, - }, - time::Duration, -}; - -use arc_swap::ArcSwap; -use cpex::cpex_core::{ - config::CpexConfig, - factory::{PluginFactory, PluginFactoryRegistry}, -}; -use rmcp::{ - ErrorData, - model::{ - CallToolRequestParams, CallToolResult, ErrorCode, GetPromptRequestParams, GetPromptResult, ReadResourceResult, - }, - serde::{Serialize, de::DeserializeOwned}, -}; -use tokio::task::JoinHandle; - -use crate::{ - config::{LoadedRuntimePluginConfig, RedisRuntimePluginConfigStore, RuntimePluginConfigStore, cpex_config}, - error::GatewayPluginRuntimeError, - hooks::{PromptPreFetchResult, RuntimeHookError, RuntimeHookState, ToolPreCallResult}, - runtime::{GatewayPluginRuntime, ResourceCallState}, -}; - -const DEFAULT_CONFIG_WATCHER_INTERVAL: Duration = Duration::from_mins(10); - -pub struct CpexRuntimeRegistry { - runtime: Arc>, - config_store: Option>, - factories: Arc, - watcher_started: AtomicBool, - watcher_interval: Duration, -} - -#[derive(Clone)] -pub struct GatewayPluginRuntimeHandle { - runtime: Arc>, -} - -struct RegistryCallState { - runtime: Arc, - state: Option, -} - -struct RegistryResourceCallState { - runtime: Arc, - state: ResourceCallState, -} - -/// Captures the resource post-hook decision and runtime for one request. -pub struct ResourceHookState { - rewritten_uri: Option, - call: Option, -} - -impl ResourceHookState { - pub fn rewritten_uri(&self) -> Option<&str> { - self.rewritten_uri.as_deref() - } - - pub async fn after_read_resource(self, response: ReadResourceResult) -> Result { - match self.call { - Some(call) => call.runtime.after_read_resource(response, call.state).await, - None => Ok(response), - } - } -} - -enum RuntimeState { - Active(Arc), - Failed(String), -} - -impl Default for CpexRuntimeRegistry { - fn default() -> Self { - Self { - runtime: Arc::new(ArcSwap::from_pointee(RuntimeState::Active(Arc::new(GatewayPluginRuntime::default())))), - config_store: None, - factories: Arc::new(PluginFactoryRegistry::new()), - watcher_started: AtomicBool::new(false), - watcher_interval: DEFAULT_CONFIG_WATCHER_INTERVAL, - } - } -} - -impl CpexRuntimeRegistry { - pub fn with_redis_config(redis_client: redis::Client) -> Self { - Self { config_store: Some(Arc::new(RedisRuntimePluginConfigStore::new(redis_client))), ..Self::default() } - } - - pub fn register_factory( - &mut self, - kind: impl Into, - factory: Box, - ) -> Result<(), GatewayPluginRuntimeError> { - let factories = Arc::get_mut(&mut self.factories).ok_or(GatewayPluginRuntimeError::FactoryRegistryShared)?; - factories.register(kind, factory); - Ok(()) - } - - pub async fn reload(&self) -> Result<(), GatewayPluginRuntimeError> { - reload_runtime(&self.runtime, self.config_store.as_ref(), &self.factories).await.map(|_| ()) - } - - pub async fn apply_config(&self, config: Option) -> Result<(), GatewayPluginRuntimeError> { - apply_runtime_config(&self.runtime, &self.factories, config).await - } - - pub fn handle(&self) -> GatewayPluginRuntimeHandle { - GatewayPluginRuntimeHandle { runtime: Arc::clone(&self.runtime) } - } - - fn start_config_watcher(&self, initial_config: Option>) -> Option> { - let config_store = self.config_store.clone()?; - if self.watcher_started.swap(true, Ordering::AcqRel) { - return None; - } - - let runtime = Arc::downgrade(&self.runtime); - let factories = Arc::clone(&self.factories); - let watcher_interval = self.watcher_interval; - Some(tokio::spawn(async move { - let mut last_applied_config = initial_config; - loop { - tokio::time::sleep(watcher_interval).await; - let Some(runtime) = runtime.upgrade() else { - break; - }; - match config_store.get_config().await { - Ok(Some(config)) => { - if last_applied_config.as_ref() == Some(&config.fingerprint) { - continue; - } - let fingerprint = config.fingerprint.clone(); - let result = match config_to_cpex(&config) { - Ok(config) => apply_runtime_config(&runtime, &factories, config).await, - Err(error) => Err(error), - }; - match result { - Ok(()) => last_applied_config = Some(fingerprint), - Err(error) => { - tracing::warn!(%error, "failed to reload CPEX runtime plugin config"); - set_runtime_failed(&runtime, &error); - last_applied_config = None; - }, - } - }, - Ok(None) => { - let error = GatewayPluginRuntimeError::ConfigMissing; - tracing::warn!(%error, "failed to reload CPEX runtime plugin config"); - set_runtime_failed(&runtime, &error); - last_applied_config = None; - }, - Err(error) => { - tracing::warn!(%error, "failed to load CPEX runtime plugin config"); - set_runtime_failed(&runtime, &error); - last_applied_config = None; - }, - } - } - })) - } -} - -async fn reload_runtime( - runtime: &ArcSwap, - config_store: Option<&Arc>, - factories: &PluginFactoryRegistry, -) -> Result>, GatewayPluginRuntimeError> { - let Some(config_store) = config_store else { - return Ok(None); - }; - match load_runtime_config(config_store, factories).await { - Ok((config, fingerprint)) => { - drop(runtime.swap(Arc::new(RuntimeState::Active(Arc::new(config))))); - Ok(fingerprint) - }, - Err(error) => { - set_runtime_failed(runtime, &error); - Err(error) - }, - } -} - -fn config_to_cpex(config: &LoadedRuntimePluginConfig) -> Result, GatewayPluginRuntimeError> { - cpex_config(&config.document).map(Some) -} - -#[cfg(test)] -impl CpexRuntimeRegistry { - fn with_config_store(config_store: Arc) -> Self { - Self { config_store: Some(config_store), ..Self::default() } - } - - fn with_config_store_interval(config_store: Arc, watcher_interval: Duration) -> Self { - Self { config_store: Some(config_store), watcher_interval, ..Self::default() } - } - - async fn before_tool_call( - &self, - request: &CallToolRequestParams, - tool_name: &str, - backend_name: &str, - ) -> Result { - self.handle().before_tool_call(request, tool_name, backend_name).await - } - - async fn after_tool_call( - &self, - tool_name: &str, - response: CallToolResult, - state: Option, - ) -> Result { - self.handle().after_tool_call(tool_name, response, state).await - } -} - -async fn apply_runtime_config( - runtime: &ArcSwap, - factories: &PluginFactoryRegistry, - config: Option, -) -> Result<(), GatewayPluginRuntimeError> { - let Some(config) = config else { - drop(runtime.swap(Arc::new(RuntimeState::Active(Arc::new(GatewayPluginRuntime::default()))))); - return Ok(()); - }; - drop( - runtime.swap(Arc::new(RuntimeState::Active(Arc::new( - GatewayPluginRuntime::from_config(config, factories).await?, - )))), - ); - Ok(()) -} - -async fn load_runtime_config( - config_store: &Arc, - factories: &PluginFactoryRegistry, -) -> Result<(GatewayPluginRuntime, Option>), GatewayPluginRuntimeError> { - let config = config_store.get_config().await?.ok_or(GatewayPluginRuntimeError::ConfigMissing)?; - let fingerprint = Some(config.fingerprint.clone()); - let runtime = match config_to_cpex(&config)? { - Some(config) => GatewayPluginRuntime::from_config(config, factories).await?, - None => GatewayPluginRuntime::default(), - }; - Ok((runtime, fingerprint)) -} - -fn set_runtime_failed(runtime: &ArcSwap, error: &GatewayPluginRuntimeError) { - drop(runtime.swap(Arc::new(RuntimeState::Failed(error.to_string())))); -} - -impl CpexRuntimeRegistry { - pub async fn initialize(&self) -> Result>, RuntimeHookError> { - let initial_config = reload_runtime(&self.runtime, self.config_store.as_ref(), &self.factories).await?; - Ok(self.start_config_watcher(initial_config)) - } -} - -impl GatewayPluginRuntimeHandle { - fn current(&self) -> Arc { - self.runtime.load_full() - } - - pub async fn before_tool_call( - &self, - request: &CallToolRequestParams, - tool_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_tool_call(request, tool_name, backend_name).await?; - if runtime.has_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 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 before_read_resource(&self, resource_uri: &str) -> Result { - let state = self.current(); - let RuntimeState::Active(runtime) = state.as_ref() else { - return Err(runtime_failed_error(state.as_ref())); - }; - let (rewritten_uri, call) = runtime.before_read_resource(resource_uri).await?; - Ok(ResourceHookState { - rewritten_uri, - call: call.map(|state| RegistryResourceCallState { runtime: Arc::clone(runtime), state }), - }) - } - - 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()) { - Some(state) => state.runtime.after_tool_call(tool_name, response, state.state.clone()).await, - None => Ok(response), - } - } - - /// Runs the tool post hooks over a streamed tool event (progress or logging - /// notification). Returns `None` when a plugin denies the event. - pub async fn after_stream_event( - &self, - tool_name: &str, - event: T, - state: Option, - ) -> Result, ErrorData> - where - T: Serialize + DeserializeOwned, - { - 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)), - } - } -} - -fn runtime_failed_error(state: &RuntimeState) -> ErrorData { - if let RuntimeState::Failed(error) = state { - tracing::warn!(%error, "rejecting MCP call because CPEX runtime is failed"); - } - ErrorData { code: ErrorCode::INTERNAL_ERROR, message: "Runtime plugin reload failed".into(), data: None } -} - -#[cfg(test)] -mod tests { - use std::{ - collections::HashMap, - sync::{ - Arc, Mutex, - atomic::{AtomicUsize, Ordering}, - }, - time::Duration, - }; - - use async_trait::async_trait; - use cpex::cpex_core::{ - cmf::{CmfHook, ContentPart, MessagePayload}, - context::PluginContext, - error::{PluginError, PluginViolation}, - factory::{PluginFactory, PluginInstance}, - hooks::{Extensions, HookHandler, PluginResult, TypedHandlerAdapter, types::cmf_hook_names}, - plugin::{Plugin, PluginConfig}, - registry::AnyHookHandler, - }; - use rmcp::model::{ - CallToolRequestParams, CallToolResult, ContentBlock, NumberOrString, ProgressNotificationParam, ProgressToken, - ReadResourceResult, ResourceContents, - }; - use serde_json::{Value, json}; - use tokio::sync::Mutex as TokioMutex; - - use contextforge_data_plane_apis::runtime_plugin_config::{ - RUNTIME_PLUGIN_CONFIG_VERSION, RuntimePluginConfigDocument, - }; - - use crate::config::LoadedRuntimePluginConfig; - 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); - const TEST_WATCHER_RETRY_COUNT: usize = 20; - const TEST_WATCHER_RETRY_INTERVAL: Duration = Duration::from_millis(20); - - #[derive(Clone, Default)] - struct MemoryConfigStore { - config: Arc>>, - calls: Arc, - } - - impl MemoryConfigStore { - fn with_config(config: RuntimePluginConfigDocument) -> Self { - Self { config: Arc::new(TokioMutex::new(Some(config))), calls: Arc::new(AtomicUsize::new(0)) } - } - - async fn set_config(&self, config: RuntimePluginConfigDocument) { - *self.config.lock().await = Some(config); - } - - async fn clear_config(&self) { - *self.config.lock().await = None; - } - - fn calls(&self) -> usize { - self.calls.load(Ordering::SeqCst) - } - } - - #[async_trait] - impl RuntimePluginConfigStore for MemoryConfigStore { - async fn get_config(&self) -> Result, GatewayPluginRuntimeError> { - self.calls.fetch_add(1, Ordering::SeqCst); - Ok(self.config.lock().await.as_ref().map(loaded_config)) - } - } - - #[derive(Default)] - struct Observations { - pre_calls: usize, - post_calls: usize, - shutdown_calls: usize, - pre_tool_call_id: Option, - post_tool_call_id: Option, - } - - #[derive(Clone, Copy, Default)] - enum PreBehavior { - #[default] - Allow, - Rewrite, - SetContext, - } - - #[derive(Clone, Copy, Default)] - enum PostBehavior { - #[default] - Allow, - Rewrite, - RewriteStreamEvent, - RewriteInvalid, - Deny, - RequireContext, - } - - struct TestPlugin { - config: PluginConfig, - observations: Arc>, - pre_behavior: PreBehavior, - post_behavior: PostBehavior, - } - - impl TestPlugin { - fn new(name: &str, hooks: Vec<&'static str>) -> Self { - Self { - config: PluginConfig { - name: name.to_owned(), - kind: "test".to_owned(), - hooks: hooks.into_iter().map(str::to_owned).collect(), - ..Default::default() - }, - observations: Arc::new(Mutex::new(Observations::default())), - pre_behavior: PreBehavior::Allow, - post_behavior: PostBehavior::Allow, - } - } - - fn rewrite_from_config(config: PluginConfig) -> Self { - Self { config, ..Self::new("generic-pre", vec![cmf_hook_names::TOOL_PRE_INVOKE]).with_pre_rewrite() } - } - - fn with_pre_rewrite(mut self) -> Self { - self.pre_behavior = PreBehavior::Rewrite; - self - } - - fn with_post_rewrite(mut self) -> Self { - self.post_behavior = PostBehavior::Rewrite; - self - } - - fn with_stream_event_rewrite(mut self) -> Self { - self.post_behavior = PostBehavior::RewriteStreamEvent; - self - } - - fn with_invalid_stream_rewrite(mut self) -> Self { - self.post_behavior = PostBehavior::RewriteInvalid; - self - } - - fn with_post_deny(mut self) -> Self { - self.post_behavior = PostBehavior::Deny; - self - } - - fn with_context_roundtrip(mut self) -> Self { - self.pre_behavior = PreBehavior::SetContext; - self.post_behavior = PostBehavior::RequireContext; - self - } - - fn observations(&self) -> Arc> { - Arc::clone(&self.observations) - } - } - - #[async_trait] - impl Plugin for TestPlugin { - fn config(&self) -> &PluginConfig { - &self.config - } - - async fn shutdown(&self) -> Result<(), Box> { - self.observations.lock().expect("observations lock poisoned").shutdown_calls += 1; - Ok(()) - } - } - - #[allow(clippy::unused_async_trait_impl)] - impl HookHandler for TestPlugin { - async fn handle( - &self, - payload: &MessagePayload, - _extensions: &Extensions, - ctx: &mut PluginContext, - ) -> PluginResult { - let is_post = payload.message.content.iter().any(|part| { - matches!( - part, - ContentPart::ToolResult { .. } | ContentPart::PromptResult { .. } | ContentPart::Resource { .. } - ) - }); - let mut observations = self.observations.lock().expect("observations lock poisoned"); - if is_post { - observations.post_calls += 1; - observations.post_tool_call_id = - payload.message.get_tool_results().first().map(|result| result.tool_call_id.clone()); - } else { - observations.pre_calls += 1; - observations.pre_tool_call_id = - payload.message.get_tool_calls().first().map(|call| call.tool_call_id.clone()); - } - drop(observations); - - if is_post { - match self.post_behavior { - PostBehavior::Allow => PluginResult::allow(), - PostBehavior::Rewrite => PluginResult::modify_payload(payload.clone()), - PostBehavior::RewriteStreamEvent => { - let mut modified = payload.clone(); - if let Some(ContentPart::ToolResult { content }) = modified - .message - .content - .iter_mut() - .find(|part| matches!(part, ContentPart::ToolResult { .. })) - && let Ok(mut progress) = - serde_json::from_value::(content.content.clone()) - { - progress.message = progress.message.map(|message| format!("plugin:{message}")); - content.content = serde_json::to_value(progress).expect("progress serializes"); - } - PluginResult::modify_payload(modified) - }, - PostBehavior::RewriteInvalid => { - let mut modified = payload.clone(); - if let Some(ContentPart::ToolResult { content }) = modified - .message - .content - .iter_mut() - .find(|part| matches!(part, ContentPart::ToolResult { .. })) - { - content.content = json!("not-a-stream-event"); - } - PluginResult::modify_payload(modified) - }, - PostBehavior::Deny => PluginResult::deny(PluginViolation::new("post_denied", "post denied")), - PostBehavior::RequireContext => { - if ctx.get_global("pre_seen") == Some(&json!(true)) { - PluginResult::allow() - } else { - PluginResult::deny( - PluginViolation::new("missing_context", "pre context missing") - .with_proto_error_code(TEST_MISSING_CONTEXT_ERROR_CODE), - ) - } - }, - } - } else { - match self.pre_behavior { - PreBehavior::Allow => PluginResult::allow(), - PreBehavior::Rewrite => { - let mut modified = payload.clone(); - if let Some(ContentPart::ToolCall { content }) = modified - .message - .content - .iter_mut() - .find(|part| matches!(part, ContentPart::ToolCall { .. })) - { - content.arguments = HashMap::from([ - ("a".to_owned(), json!(TEST_REWRITTEN_SUM_A)), - ("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 => { - ctx.set_global("pre_seen", json!(true)); - PluginResult::allow() - }, - } - } - } - } - - struct TestPluginFactory { - observations: Arc>, - pre_behavior: PreBehavior, - post_behavior: PostBehavior, - } - - impl TestPluginFactory { - fn from_plugin(plugin: &TestPlugin) -> Self { - Self { - observations: Arc::clone(&plugin.observations), - pre_behavior: plugin.pre_behavior, - post_behavior: plugin.post_behavior, - } - } - } - - impl PluginFactory for TestPluginFactory { - fn create(&self, config: &PluginConfig) -> Result> { - let plugin = Arc::new(TestPlugin { - config: config.clone(), - observations: Arc::clone(&self.observations), - pre_behavior: self.pre_behavior, - post_behavior: self.post_behavior, - }); - let handlers = config - .hooks - .iter() - .filter_map(|hook| { - let hook = match hook.as_str() { - cmf_hook_names::TOOL_PRE_INVOKE => cmf_hook_names::TOOL_PRE_INVOKE, - cmf_hook_names::TOOL_POST_INVOKE => cmf_hook_names::TOOL_POST_INVOKE, - cmf_hook_names::RESOURCE_PRE_FETCH => cmf_hook_names::RESOURCE_PRE_FETCH, - cmf_hook_names::RESOURCE_POST_FETCH => cmf_hook_names::RESOURCE_POST_FETCH, - _ => return None, - }; - Some(( - hook, - Arc::new(TypedHandlerAdapter::::new(Arc::clone(&plugin))) - as Arc, - )) - }) - .collect(); - let plugin: Arc = plugin; - Ok(PluginInstance { plugin, handlers }) - } - } - - fn sum_request(a: i64, b: i64) -> CallToolRequestParams { - CallToolRequestParams::new("sum") - .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") - } - - fn config_document(cpex: Value) -> RuntimePluginConfigDocument { - RuntimePluginConfigDocument { - version: RUNTIME_PLUGIN_CONFIG_VERSION, - cpex: serde_json::from_value(cpex).expect("test CPEX config parses"), - } - } - - fn loaded_config(document: &RuntimePluginConfigDocument) -> LoadedRuntimePluginConfig { - LoadedRuntimePluginConfig::decode(serde_json::to_vec(document).expect("test CPEX config serializes")) - .expect("test CPEX config decodes") - } - - fn plugin_config(plugins: &[Arc]) -> RuntimePluginConfigDocument { - config_document(json!({ - "plugins": plugins.iter().map(|plugin| { - json!({ - "name": plugin.config.name.clone(), - "kind": plugin.config.kind.clone(), - "hooks": plugin.config.hooks.clone(), - }) - }).collect::>() - })) - } - - fn expect_runtime_failed(result: Result) -> ErrorData { - match result { - Ok(_) => panic!("runtime should be failed"), - Err(error) => error, - } - } - - async fn runtime_with_plugin(plugin: &Arc, config: RuntimePluginConfigDocument) -> CpexRuntimeRegistry { - let mut runtime = CpexRuntimeRegistry::with_config_store(Arc::new(MemoryConfigStore::with_config(config))); - runtime - .register_factory("test", Box::new(TestPluginFactory::from_plugin(plugin))) - .expect("test factory registers"); - runtime.initialize().await.expect("runtime initializes"); - runtime - } - - #[tokio::test(flavor = "multi_thread", worker_threads = 1)] - async fn runtime_config_store_is_loaded_on_initialize() { - let config_store = MemoryConfigStore::with_config(config_document(json!({ "plugins": [] }))); - let runtime = CpexRuntimeRegistry::with_config_store(Arc::new(config_store.clone())); - - let handle = runtime.initialize().await.expect("runtime initializes"); - - assert!(handle.is_some()); - assert!(config_store.calls() >= 1); - } - - #[tokio::test(flavor = "multi_thread", worker_threads = 1)] - async fn missing_runtime_plugin_config_is_rejected_on_initialize() { - let runtime = CpexRuntimeRegistry::with_config_store(Arc::new(MemoryConfigStore::default())); - - let error = runtime.initialize().await.expect_err("missing config is rejected"); - - assert_eq!("runtime plugin config is missing", error.to_string()); - } - - #[tokio::test(flavor = "multi_thread", worker_threads = 1)] - async fn invalid_runtime_plugin_config_documents_are_rejected() { - for config in [RuntimePluginConfigDocument { version: 2, cpex: CpexConfig::default() }] { - let runtime = CpexRuntimeRegistry::with_config_store(Arc::new(MemoryConfigStore::with_config(config))); - let error = runtime.initialize().await.expect_err("invalid config is rejected"); - - assert_eq!("runtime plugin config is in wrong format", error.to_string()); - } - } - - #[tokio::test(flavor = "multi_thread", worker_threads = 1)] - async fn unsupported_runtime_plugin_config_is_rejected() { - for cpex in [ - json!({ "plugin_settings": { "routing_enabled": true }, "plugins": [] }), - json!({ "plugin_settings": { "fail_on_plugin_error": true }, "plugins": [] }), - json!({ - "plugins": [{ - "name": "scoped", - "kind": "test", - "hooks": [cmf_hook_names::TOOL_PRE_INVOKE], - "conditions": [{ "tools": ["sum"] }] - }] - }), - json!({ "plugins": [{ "name": "llm", "kind": "test", "hooks": [cmf_hook_names::LLM_INPUT] }] }), - ] { - let runtime = - CpexRuntimeRegistry::with_config_store(Arc::new(MemoryConfigStore::with_config(config_document(cpex)))); - let error = runtime.initialize().await.expect_err("unsupported config is rejected"); - - assert_eq!("runtime plugin config is unsupported", error.to_string()); - } - } - - #[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 resource_pre_hook_runs_for_a_canonical_uri() { - let plugin = Arc::new(TestPlugin::new("resource", vec![cmf_hook_names::RESOURCE_PRE_FETCH])); - let observations = plugin.observations(); - let runtime = runtime_with_plugin(&plugin, plugin_config(&[Arc::clone(&plugin)])).await; - - runtime.handle().before_read_resource("file:///password.env").await.expect("resource pre hook runs"); - - assert_eq!(1, observations.lock().expect("observations lock poisoned").pre_calls); - } - - #[tokio::test(flavor = "multi_thread", worker_threads = 1)] - async fn resource_without_post_hook_keeps_its_decision_across_reload() { - let plugin = Arc::new(TestPlugin::new("resource", vec![cmf_hook_names::RESOURCE_POST_FETCH])); - let observations = plugin.observations(); - let runtime = runtime_with_plugin(&plugin, plugin_config(&[Arc::clone(&plugin)])).await; - runtime.apply_config(None).await.expect("disable hooks"); - let state = runtime.handle().before_read_resource("file:///password.env").await.expect("request starts"); - assert!(state.call.is_none(), "no post-hook state allocation"); - runtime.apply_config(Some(plugin_config(&[plugin]).cpex)).await.expect("enable hooks"); - let response = ReadResourceResult::new(vec![ResourceContents::text("original", "file:///password.env")]); - state.after_read_resource(response).await.expect("in-flight decision survives reload"); - assert_eq!(0, observations.lock().expect("observations lock poisoned").post_calls); - } - - #[tokio::test(flavor = "multi_thread", worker_threads = 1)] - async fn resource_post_hook_keeps_its_runtime_across_reload() { - let plugin = Arc::new(TestPlugin::new("resource", vec![cmf_hook_names::RESOURCE_POST_FETCH]).with_post_deny()); - let runtime = runtime_with_plugin(&plugin, plugin_config(&[Arc::clone(&plugin)])).await; - let state = runtime.handle().before_read_resource("file:///password.env").await.expect("request starts"); - runtime.apply_config(None).await.expect("disable hooks"); - let response = ReadResourceResult::new(vec![ResourceContents::text("secret", "file:///password.env")]); - let error = state.after_read_resource(response).await.expect_err("captured policy still denies"); - assert_eq!("Plugin denied resource", error.message); - } - - #[tokio::test(flavor = "multi_thread", worker_threads = 1)] - async fn resource_hooks_preserve_context_across_the_backend_call() { - let plugin = Arc::new( - TestPlugin::new("resource", vec![cmf_hook_names::RESOURCE_PRE_FETCH, cmf_hook_names::RESOURCE_POST_FETCH]) - .with_context_roundtrip(), - ); - let observations = plugin.observations(); - let runtime = runtime_with_plugin(&plugin, plugin_config(&[Arc::clone(&plugin)])).await; - let pre = runtime.handle().before_read_resource("file:///password.env").await.expect("resource pre hook runs"); - let response = ReadResourceResult::new(vec![ResourceContents::text("secret", "file:///password.env")]); - - pre.after_read_resource(response).await.expect("resource post hook receives pre context"); - - let observations = observations.lock().expect("observations lock poisoned"); - assert_eq!(1, observations.pre_calls); - assert_eq!(1, observations.post_calls); - } - - #[tokio::test(flavor = "multi_thread", worker_threads = 1)] - async fn runtime_config_loads_registered_factory_plugin() { - let plugin = - Arc::new(TestPlugin::new("configured-pre", vec![cmf_hook_names::TOOL_PRE_INVOKE]).with_pre_rewrite()); - let observations = plugin.observations(); - let runtime = runtime_with_plugin(&plugin, plugin_config(&[Arc::clone(&plugin)])).await; - - let result = runtime.before_tool_call(&sum_request(1, 2), "sum", "backend").await.expect("pre hook runs"); - - assert!(matches!(result.arguments, ToolArgumentsUpdate::Replace(Some(_)))); - assert_eq!(1, observations.lock().expect("observations lock poisoned").pre_calls); - } - - #[tokio::test(flavor = "multi_thread", worker_threads = 1)] - async fn runtime_config_loads_generic_cmf_factory_plugin() { - let config = config_document(json!({ - "plugins": [{ - "name": "generic-pre", - "kind": "generic", - "hooks": [cmf_hook_names::TOOL_PRE_INVOKE] - }] - })); - 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.before_tool_call(&sum_request(1, 2), "sum", "backend").await.expect("pre hook runs"); - - 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 = - Arc::new(TestPlugin::new("configured-pre", vec![cmf_hook_names::TOOL_PRE_INVOKE]).with_pre_rewrite()); - let observations = plugin.observations(); - let config_store = MemoryConfigStore::with_config(config_document(json!({ "plugins": [] }))); - let mut runtime = CpexRuntimeRegistry::with_config_store(Arc::new(config_store.clone())); - runtime - .register_factory("test", Box::new(TestPluginFactory::from_plugin(&plugin))) - .expect("test factory registers"); - runtime.initialize().await.expect("runtime initializes"); - - let result = runtime.before_tool_call(&sum_request(1, 2), "sum", "backend").await.expect("pre hook skips"); - assert!(matches!(result.arguments, ToolArgumentsUpdate::Unchanged)); - - config_store.set_config(plugin_config(&[Arc::clone(&plugin)])).await; - runtime.reload().await.expect("runtime reloads"); - let result = runtime.before_tool_call(&sum_request(1, 2), "sum", "backend").await.expect("pre hook runs"); - assert!(matches!(result.arguments, ToolArgumentsUpdate::Replace(Some(_)))); - - config_store.set_config(config_document(json!({ "plugins": [] }))).await; - runtime.reload().await.expect("runtime reloads"); - let result = runtime.before_tool_call(&sum_request(1, 2), "sum", "backend").await.expect("pre hook skips"); - assert!(matches!(result.arguments, ToolArgumentsUpdate::Unchanged)); - assert_eq!(1, observations.lock().expect("observations lock poisoned").pre_calls); - } - - #[tokio::test(flavor = "multi_thread", worker_threads = 1)] - async fn failed_runtime_reload_rejects_new_calls_until_valid_reload() { - let plugin = - Arc::new(TestPlugin::new("configured-pre", vec![cmf_hook_names::TOOL_PRE_INVOKE]).with_pre_rewrite()); - let observations = plugin.observations(); - let config_store = MemoryConfigStore::with_config(plugin_config(&[Arc::clone(&plugin)])); - let mut runtime = CpexRuntimeRegistry::with_config_store(Arc::new(config_store.clone())); - runtime - .register_factory("test", Box::new(TestPluginFactory::from_plugin(&plugin))) - .expect("test factory registers"); - runtime.initialize().await.expect("runtime initializes"); - - config_store.set_config(RuntimePluginConfigDocument { version: 2, cpex: CpexConfig::default() }).await; - runtime.reload().await.expect_err("invalid reload fails"); - let error = expect_runtime_failed(runtime.before_tool_call(&sum_request(1, 2), "sum", "backend").await); - assert_eq!(ErrorCode::INTERNAL_ERROR, error.code); - assert_eq!("Runtime plugin reload failed", error.message); - - config_store.clear_config().await; - runtime.reload().await.expect_err("missing reload fails"); - let error = expect_runtime_failed(runtime.before_tool_call(&sum_request(1, 2), "sum", "backend").await); - assert_eq!(ErrorCode::INTERNAL_ERROR, error.code); - assert_eq!("Runtime plugin reload failed", error.message); - - config_store.set_config(plugin_config(&[Arc::clone(&plugin)])).await; - runtime.reload().await.expect("runtime recovers"); - let result = runtime.before_tool_call(&sum_request(1, 2), "sum", "backend").await.expect("pre hook runs"); - assert!(matches!(result.arguments, ToolArgumentsUpdate::Replace(Some(_)))); - assert_eq!(1, observations.lock().expect("observations lock poisoned").pre_calls); - } - - #[tokio::test(flavor = "multi_thread", worker_threads = 1)] - async fn combined_plugin_preserves_context_from_pre_to_post_across_replacement() { - let plugin = Arc::new( - TestPlugin::new("context", vec![cmf_hook_names::TOOL_PRE_INVOKE, cmf_hook_names::TOOL_POST_INVOKE]) - .with_context_roundtrip(), - ); - let observations = plugin.observations(); - let config_store = MemoryConfigStore::with_config(plugin_config(&[Arc::clone(&plugin)])); - let mut runtime = CpexRuntimeRegistry::with_config_store(Arc::new(config_store.clone())); - runtime - .register_factory("test", Box::new(TestPluginFactory::from_plugin(&plugin))) - .expect("test factory registers"); - runtime.initialize().await.expect("runtime initializes"); - - let pre = runtime.before_tool_call(&sum_request(1, 2), "sum", "backend").await.expect("pre hook runs"); - config_store.set_config(config_document(json!({ "plugins": [] }))).await; - runtime.reload().await.expect("runtime reloads"); - let response = CallToolResult::success(vec![ContentBlock::text("3")]); - runtime.after_tool_call("sum", response, pre.state).await.expect("post hook runs"); - - let observations = observations.lock().expect("observations lock poisoned"); - assert_eq!(1, observations.post_calls); - assert_eq!(observations.pre_tool_call_id, observations.post_tool_call_id); - } - - #[tokio::test(flavor = "multi_thread", worker_threads = 1)] - async fn post_only_runtime_does_not_apply_new_post_hook_to_in_flight_call() { - let plugin = Arc::new(TestPlugin::new("post", vec![cmf_hook_names::TOOL_POST_INVOKE]).with_post_rewrite()); - let observations = plugin.observations(); - let config_store = MemoryConfigStore::with_config(config_document(json!({ "plugins": [] }))); - let mut runtime = CpexRuntimeRegistry::with_config_store(Arc::new(config_store.clone())); - runtime - .register_factory("test", Box::new(TestPluginFactory::from_plugin(&plugin))) - .expect("test factory registers"); - runtime.initialize().await.expect("runtime initializes"); - - let pre = runtime.before_tool_call(&sum_request(1, 2), "sum", "backend").await.expect("pre hook skips"); - config_store.set_config(plugin_config(&[Arc::clone(&plugin)])).await; - runtime.reload().await.expect("runtime reloads"); - let response = CallToolResult::success(vec![ContentBlock::text("3")]); - runtime.after_tool_call("sum", response, pre.state).await.expect("post hook skips"); - - assert_eq!(0, observations.lock().expect("observations lock poisoned").post_calls); - } - - #[tokio::test(flavor = "multi_thread", worker_threads = 1)] - async fn stream_event_without_state_skips_post_hook() { - let plugin = - Arc::new(TestPlugin::new("post", vec![cmf_hook_names::TOOL_POST_INVOKE]).with_stream_event_rewrite()); - let observations = plugin.observations(); - let runtime = runtime_with_plugin(&plugin, plugin_config(&[Arc::clone(&plugin)])).await; - - let event = runtime.handle().after_stream_event("sum", progress_event(), None).await.expect("event passes"); - - assert_eq!(Some("step 1/2"), event.expect("event is kept").message.as_deref()); - assert_eq!(0, observations.lock().expect("observations lock poisoned").post_calls); - } - - #[tokio::test(flavor = "multi_thread", worker_threads = 1)] - async fn stream_event_is_rewritten_by_post_hook() { - let plugin = - Arc::new(TestPlugin::new("post", vec![cmf_hook_names::TOOL_POST_INVOKE]).with_stream_event_rewrite()); - let observations = plugin.observations(); - let runtime = runtime_with_plugin(&plugin, plugin_config(&[Arc::clone(&plugin)])).await; - - let pre = runtime.before_tool_call(&sum_request(1, 2), "sum", "backend").await.expect("pre state is created"); - let event = - runtime.handle().after_stream_event("sum", progress_event(), pre.state).await.expect("event passes"); - - assert_eq!(Some("plugin:step 1/2"), event.expect("event is kept").message.as_deref()); - assert_eq!(1, observations.lock().expect("observations lock poisoned").post_calls); - } - - #[tokio::test(flavor = "multi_thread", worker_threads = 1)] - async fn denied_stream_event_is_dropped() { - let plugin = Arc::new(TestPlugin::new("post", vec![cmf_hook_names::TOOL_POST_INVOKE]).with_post_deny()); - let runtime = runtime_with_plugin(&plugin, plugin_config(&[Arc::clone(&plugin)])).await; - - let pre = runtime.before_tool_call(&sum_request(1, 2), "sum", "backend").await.expect("pre state is created"); - let event = runtime - .handle() - .after_stream_event("sum", progress_event(), pre.state) - .await - .expect("deny drops the event"); - - assert!(event.is_none()); - } - - #[tokio::test(flavor = "multi_thread", worker_threads = 1)] - async fn invalid_stream_event_rewrite_is_rejected() { - let plugin = - Arc::new(TestPlugin::new("post", vec![cmf_hook_names::TOOL_POST_INVOKE]).with_invalid_stream_rewrite()); - let runtime = runtime_with_plugin(&plugin, plugin_config(&[Arc::clone(&plugin)])).await; - - let pre = runtime.before_tool_call(&sum_request(1, 2), "sum", "backend").await.expect("pre state is created"); - let error = runtime - .handle() - .after_stream_event("sum", progress_event(), pre.state) - .await - .expect_err("invalid rewrite is rejected"); - - assert_eq!(ErrorCode::INVALID_PARAMS, error.code); - } - - #[tokio::test(flavor = "multi_thread", worker_threads = 1)] - async fn replaced_runtime_shutdowns_on_drop() { - let plugin = Arc::new(TestPlugin::new("pre", vec![cmf_hook_names::TOOL_PRE_INVOKE]).with_pre_rewrite()); - let observations = plugin.observations(); - let config_store = MemoryConfigStore::with_config(plugin_config(&[Arc::clone(&plugin)])); - let mut runtime = CpexRuntimeRegistry::with_config_store(Arc::new(config_store.clone())); - runtime - .register_factory("test", Box::new(TestPluginFactory::from_plugin(&plugin))) - .expect("test factory registers"); - runtime.initialize().await.expect("runtime initializes"); - - config_store.set_config(config_document(json!({ "plugins": [] }))).await; - runtime.reload().await.expect("runtime reloads"); - - for _ in 0..TEST_SHUTDOWN_RETRY_COUNT { - if observations.lock().expect("observations lock poisoned").shutdown_calls > 0 { - return; - } - tokio::time::sleep(TEST_SHUTDOWN_RETRY_INTERVAL).await; - } - panic!("replaced runtime did not shut down"); - } - - #[tokio::test(flavor = "multi_thread", worker_threads = 1)] - async fn watcher_applies_config_changes() { - let plugin = - Arc::new(TestPlugin::new("configured-pre", vec![cmf_hook_names::TOOL_PRE_INVOKE]).with_pre_rewrite()); - let observations = plugin.observations(); - let config_store = MemoryConfigStore::with_config(config_document(json!({ "plugins": [] }))); - let mut runtime = - CpexRuntimeRegistry::with_config_store_interval(Arc::new(config_store.clone()), TEST_WATCHER_INTERVAL); - runtime - .register_factory("test", Box::new(TestPluginFactory::from_plugin(&plugin))) - .expect("test factory registers"); - let handle = runtime.initialize().await.expect("runtime initializes"); - assert!(handle.is_some()); - - config_store.set_config(plugin_config(&[Arc::clone(&plugin)])).await; - for _ in 0..TEST_WATCHER_RETRY_COUNT { - let result = runtime.before_tool_call(&sum_request(1, 2), "sum", "backend").await.expect("pre hook runs"); - if matches!(result.arguments, ToolArgumentsUpdate::Replace(Some(_))) { - config_store.clear_config().await; - tokio::time::sleep(TEST_WATCHER_INTERVAL + TEST_WATCHER_RETRY_INTERVAL).await; - let error = expect_runtime_failed(runtime.before_tool_call(&sum_request(1, 2), "sum", "backend").await); - assert_eq!(ErrorCode::INTERNAL_ERROR, error.code); - - config_store.set_config(plugin_config(&[Arc::clone(&plugin)])).await; - for _ in 0..TEST_WATCHER_RETRY_COUNT { - if let Ok(result) = runtime.before_tool_call(&sum_request(1, 2), "sum", "backend").await - && matches!(result.arguments, ToolArgumentsUpdate::Replace(Some(_))) - { - assert_eq!(2, observations.lock().expect("observations lock poisoned").pre_calls); - return; - } - tokio::time::sleep(TEST_WATCHER_RETRY_INTERVAL).await; - } - panic!("config watcher did not recover from missing plugin config"); - } - tokio::time::sleep(TEST_WATCHER_RETRY_INTERVAL).await; - } - panic!("config watcher did not apply plugin config"); - } - - #[tokio::test(flavor = "multi_thread", worker_threads = 1)] - async fn initialize_without_config_store_returns_no_watcher() { - let runtime = CpexRuntimeRegistry::default(); - - let handle = runtime.initialize().await.expect("runtime initializes"); - - assert!(handle.is_none()); - } -} diff --git a/crates/contextforge-data-plane-cpex/src/hooks.rs b/crates/contextforge-data-plane-cpex/src/hooks.rs index b6a9a63a..9291c391 100644 --- a/crates/contextforge-data-plane-cpex/src/hooks.rs +++ b/crates/contextforge-data-plane-cpex/src/hooks.rs @@ -1,59 +1,38 @@ -use std::{any::Any, sync::Arc}; - -use rmcp::model::{CallToolRequestParams, GetPromptRequestParams}; use serde_json::{Map, Value}; pub type RuntimeHookError = Box; -pub type RuntimeHookState = Arc; -#[derive(Debug)] -pub enum ToolArgumentsUpdate { +#[derive(Debug, Default)] +pub enum ArgumentsUpdate { + #[default] Unchanged, Replace(Option>), } -impl ToolArgumentsUpdate { - pub fn apply_to_request(self, request: &mut CallToolRequestParams, routed_tool_name: &str) { - request.name = routed_tool_name.to_owned().into(); - if let Self::Replace(arguments) = self { - request.arguments = arguments; +impl ArgumentsUpdate { + pub(crate) fn from_modified(original: Option<&Map>, modified: Map) -> Self { + if original == Some(&modified) || (original.is_none() && modified.is_empty()) { + Self::Unchanged + } else { + Self::Replace(Some(modified)) } } -} - -pub struct ToolPreCallResult { - pub arguments: ToolArgumentsUpdate, - pub state: Option, -} - -impl ToolPreCallResult { - pub fn unchanged() -> Self { - 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 fn apply_to(self, arguments: &mut Option>) { + if let Self::Replace(replacement) = self { + *arguments = replacement; } } } -pub struct PromptPreFetchResult { - pub arguments: PromptArgumentsUpdate, - pub state: Option, +/// Argument edits and the typed post-hook state captured before backend I/O. +pub struct PreHookResult { + pub arguments: ArgumentsUpdate, + pub state: Option, } -impl PromptPreFetchResult { - pub fn unchanged() -> Self { - Self { arguments: PromptArgumentsUpdate::Unchanged, state: None } +impl Default for PreHookResult { + fn default() -> Self { + Self { arguments: ArgumentsUpdate::Unchanged, state: None } } } diff --git a/crates/contextforge-data-plane-cpex/src/lib.rs b/crates/contextforge-data-plane-cpex/src/lib.rs index 3d024272..5613e4cf 100644 --- a/crates/contextforge-data-plane-cpex/src/lib.rs +++ b/crates/contextforge-data-plane-cpex/src/lib.rs @@ -2,15 +2,17 @@ mod cmf; mod config; mod error; mod factory; -mod handle; mod hooks; -mod pipeline; +mod prompts; +mod registry; +mod resources; mod runtime; +mod tools; pub use error::GatewayPluginRuntimeError; pub use factory::CmfPluginFactory; -pub use handle::{CpexRuntimeRegistry, GatewayPluginRuntimeHandle, ResourceHookState}; -pub use hooks::{ - PromptArgumentsUpdate, PromptPreFetchResult, RuntimeHookError, RuntimeHookState, ToolArgumentsUpdate, - ToolPreCallResult, -}; +pub use hooks::{ArgumentsUpdate, PreHookResult, RuntimeHookError}; +pub use prompts::PromptHookState; +pub use registry::{CpexRuntimeRegistry, GatewayPluginRuntimeHandle}; +pub use resources::ResourceHookState; +pub use tools::ToolHookState; diff --git a/crates/contextforge-data-plane-cpex/src/pipeline.rs b/crates/contextforge-data-plane-cpex/src/pipeline.rs deleted file mode 100644 index 93abb715..00000000 --- a/crates/contextforge-data-plane-cpex/src/pipeline.rs +++ /dev/null @@ -1,162 +0,0 @@ -use cpex::cpex_core::cmf::MessagePayload; -use cpex::cpex_core::executor::PipelineResult; -use rmcp::{ - ErrorData, - model::{CallToolResult, ErrorCode, GetPromptResult, ReadResourceResult}, - serde::de::DeserializeOwned, -}; -use tracing::warn; - -use crate::{ - PromptArgumentsUpdate, ToolArgumentsUpdate, - cmf::{ - prompt_request_arguments, prompt_result_rejection, prompt_result_response, resource_result_response, - tool_call_arguments, tool_result_content, tool_result_response, - }, -}; - -pub(crate) fn modified_message_payload(result: &PipelineResult) -> Option<&MessagePayload> { - result.modified_payload.as_ref().and_then(|payload| payload.as_any().downcast_ref::()) -} - -pub(crate) fn effective_pre_args( - original_args: Option<&serde_json::Map>, - pre_result: &PipelineResult, -) -> Result { - let Some(modified_payload) = modified_message_payload(pre_result) else { - return Ok(ToolArgumentsUpdate::Unchanged); - }; - - let Some(arguments) = tool_call_arguments(modified_payload) else { - return Err(ErrorData { - code: ErrorCode::INVALID_PARAMS, - message: "Plugin modified tool payload without a tool call".into(), - data: None, - }); - }; - - if original_args == Some(&arguments) || (original_args.is_none() && arguments.is_empty()) { - Ok(ToolArgumentsUpdate::Unchanged) - } else { - Ok(ToolArgumentsUpdate::Replace(Some(arguments.clone()))) - } -} - -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), - None => original, - } -} - -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_pre_resource_uri(result: &PipelineResult) -> Result, ErrorData> { - let Some(payload) = modified_message_payload(result) else { return Ok(None) }; - let [cpex::cpex_core::cmf::ContentPart::ResourceRef { content }] = payload.message.content.as_slice() else { - return Err(ErrorData::internal_error("Plugin returned an invalid resource request", None)); - }; - Ok(Some(content.uri.clone())) -} - -pub(crate) fn effective_post_resource_result( - original: ReadResourceResult, - result: &PipelineResult, -) -> Result { - let Some(payload) = modified_message_payload(result) else { - return Ok(original); - }; - resource_result_response(original, payload) - .ok_or_else(|| ErrorData::internal_error("Plugin returned a resource result the gateway cannot apply", None)) -} - -pub(crate) fn effective_post_json(original: T, result: &PipelineResult) -> Result -where - T: DeserializeOwned, -{ - let Some(payload) = modified_message_payload(result) else { - return Ok(original); - }; - let Some(content) = tool_result_content(payload) else { - return Err(ErrorData { - code: ErrorCode::INVALID_PARAMS, - message: "Plugin modified stream event payload without a tool result".into(), - data: None, - }); - }; - serde_json::from_value::(content).map_err(|error| ErrorData { - code: ErrorCode::INVALID_PARAMS, - message: format!("Plugin modified stream event payload with invalid JSON: {error}").into(), - data: None, - }) -} - -pub(crate) fn plugin_denied_error(subject: &str, result: PipelineResult) -> ErrorData { - let code = result - .violation - .and_then(|violation| { - 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: format!("Plugin denied {subject}").into(), data: None } -} - -pub(crate) fn log_pipeline_errors(hook: &'static str, result: &PipelineResult) { - for error in &result.errors { - warn!( - hook, - plugin = error.plugin_name, - code = error.code.as_deref().unwrap_or(""), - proto_error_code = error.proto_error_code, - "CPEX plugin soft error" - ); - } -} diff --git a/crates/contextforge-data-plane-cpex/src/prompts/mod.rs b/crates/contextforge-data-plane-cpex/src/prompts/mod.rs new file mode 100644 index 00000000..74fbe51e --- /dev/null +++ b/crates/contextforge-data-plane-cpex/src/prompts/mod.rs @@ -0,0 +1,287 @@ +use std::collections::HashMap; + +use cpex::cpex_core::cmf::{ + AudioSource, ContentPart, ImageSource, Message, MessagePayload, PromptRequest, PromptResult, + Resource as CmfResource, ResourceReference, ResourceType, Role, +}; +use rmcp::{ + ErrorData, + model::{ + ContentBlock, GetPromptRequestParams, GetPromptResult, PromptMessage, Resource as McpResource, + ResourceContents, Role as McpRole, + }, +}; +use serde_json::{Map, Value}; + +use crate::{ + ArgumentsUpdate, GatewayPluginRuntimeHandle, PreHookResult, + cmf::{CmfResponse, Operation, message_payload}, + runtime::CallState, +}; + +fn prompt_request_payload( + request: &GetPromptRequestParams, + prompt_name: &str, + backend_name: &str, + prompt_request_id: &str, +) -> MessagePayload { + message_payload( + Role::User, + 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()), + }, + }], + ) +} + +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::>()) +} + +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::>(); + + message_payload( + Role::Assistant, + 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, + }, + }], + ) +} + +fn prompt_result(payload: &MessagePayload) -> Option<&PromptResult> { + let results = payload.message.get_prompt_results(); + let [result] = results.as_slice() else { return None }; + Some(*result) +} + +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. +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)) +} + +/// The runtime and plugin context captured before fetching one prompt. +pub struct PromptHookState(CallState); + +impl GatewayPluginRuntimeHandle { + pub async fn before_get_prompt( + &self, + request: &GetPromptRequestParams, + prompt_name: &str, + backend_name: &str, + ) -> Result, ErrorData> { + let (arguments, state) = self + .current()? + .before( + Operation::Prompt, + prompt_name, + |id| prompt_request_payload(request, prompt_name, backend_name, id), + |payload, id| { + let arguments = + prompt_request_arguments(payload, prompt_name, backend_name, id).ok_or_else(|| { + ErrorData::invalid_params("Plugin returned a prompt request the gateway cannot apply", None) + })?; + Ok(ArgumentsUpdate::from_modified(request.arguments.as_ref(), arguments)) + }, + ) + .await?; + Ok(PreHookResult { arguments, state: state.map(PromptHookState) }) + } +} + +impl PromptHookState { + pub async fn after_get_prompt(mut self, response: GetPromptResult) -> Result { + self.0.after(response).await + } +} + +impl CmfResponse for GetPromptResult { + const OPERATION: Operation = Operation::Prompt; + + fn to_payload(&self, name: &str, id: &str) -> Result { + Ok(prompt_result_payload(self, name, id)) + } + + fn apply_payload(self, payload: &MessagePayload, name: &str, id: &str) -> Result { + if let Some(message) = prompt_result_rejection(payload) { + return Err(ErrorData::invalid_request(message, None)); + } + prompt_result_response(self, payload, name, id) + .ok_or_else(|| ErrorData::internal_error("Plugin returned a prompt result the gateway cannot apply", None)) + } +} + +#[cfg(test)] +mod tests; diff --git a/crates/contextforge-data-plane-cpex/src/prompts/tests.rs b/crates/contextforge-data-plane-cpex/src/prompts/tests.rs new file mode 100644 index 00000000..3bb0126d --- /dev/null +++ b/crates/contextforge-data-plane-cpex/src/prompts/tests.rs @@ -0,0 +1,436 @@ +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); +} diff --git a/crates/contextforge-data-plane-cpex/src/registry/mod.rs b/crates/contextforge-data-plane-cpex/src/registry/mod.rs new file mode 100644 index 00000000..180653de --- /dev/null +++ b/crates/contextforge-data-plane-cpex/src/registry/mod.rs @@ -0,0 +1,181 @@ +use std::{ + sync::{ + Arc, + atomic::{AtomicBool, Ordering}, + }, + time::Duration, +}; + +use arc_swap::ArcSwap; +use cpex::cpex_core::{ + config::CpexConfig, + factory::{PluginFactory, PluginFactoryRegistry}, +}; +use rmcp::{ErrorData, model::ErrorCode}; +use tokio::task::JoinHandle; + +use crate::{ + config::{RedisRuntimePluginConfigStore, RuntimePluginConfigStore, cpex_config}, + error::GatewayPluginRuntimeError, + hooks::RuntimeHookError, + runtime::GatewayPluginRuntime, +}; + +const DEFAULT_CONFIG_WATCHER_INTERVAL: Duration = Duration::from_mins(10); + +pub struct CpexRuntimeRegistry { + runtime: Arc>, + config_store: Option>, + factories: Arc, + watcher_started: AtomicBool, + watcher_interval: Duration, +} + +#[derive(Clone)] +pub struct GatewayPluginRuntimeHandle { + runtime: Arc>, +} + +enum RuntimeState { + Active(Arc), + Failed(String), +} + +impl Default for CpexRuntimeRegistry { + fn default() -> Self { + Self { + runtime: Arc::new(ArcSwap::from_pointee(RuntimeState::Active(Arc::new(GatewayPluginRuntime::default())))), + config_store: None, + factories: Arc::new(PluginFactoryRegistry::new()), + watcher_started: AtomicBool::new(false), + watcher_interval: DEFAULT_CONFIG_WATCHER_INTERVAL, + } + } +} + +impl CpexRuntimeRegistry { + pub fn with_redis_config(redis_client: redis::Client) -> Self { + Self { config_store: Some(Arc::new(RedisRuntimePluginConfigStore::new(redis_client))), ..Self::default() } + } + + pub fn register_factory( + &mut self, + kind: impl Into, + factory: Box, + ) -> Result<(), GatewayPluginRuntimeError> { + let factories = Arc::get_mut(&mut self.factories).ok_or(GatewayPluginRuntimeError::FactoryRegistryShared)?; + factories.register(kind, factory); + Ok(()) + } + + pub async fn reload(&self) -> Result<(), GatewayPluginRuntimeError> { + reload_runtime(&self.runtime, self.config_store.as_ref(), &self.factories, None).await.map(|_| ()) + } + + pub async fn apply_config(&self, config: Option) -> Result<(), GatewayPluginRuntimeError> { + apply_runtime_config(&self.runtime, &self.factories, config).await + } + + pub fn handle(&self) -> GatewayPluginRuntimeHandle { + GatewayPluginRuntimeHandle { runtime: Arc::clone(&self.runtime) } + } + + fn start_config_watcher(&self, initial_config: Option>) -> Option> { + let config_store = self.config_store.clone()?; + if self.watcher_started.swap(true, Ordering::AcqRel) { + return None; + } + + let runtime = Arc::downgrade(&self.runtime); + let factories = Arc::clone(&self.factories); + let watcher_interval = self.watcher_interval; + Some(tokio::spawn(async move { + let mut last_applied_config = initial_config; + loop { + tokio::time::sleep(watcher_interval).await; + let Some(runtime) = runtime.upgrade() else { + break; + }; + match reload_runtime(&runtime, Some(&config_store), &factories, last_applied_config.as_deref()).await { + Ok(Some(fingerprint)) => last_applied_config = Some(fingerprint), + Ok(None) => {}, + Err(error) => { + tracing::warn!(%error, "failed to reload CPEX runtime plugin config"); + last_applied_config = None; + }, + } + } + })) + } +} + +async fn reload_runtime( + runtime: &ArcSwap, + config_store: Option<&Arc>, + factories: &PluginFactoryRegistry, + last_applied_config: Option<&[u8]>, +) -> Result>, GatewayPluginRuntimeError> { + let Some(config_store) = config_store else { + return Ok(None); + }; + let result = async { + let config = config_store.get_config().await?.ok_or(GatewayPluginRuntimeError::ConfigMissing)?; + if last_applied_config == Some(config.fingerprint.as_slice()) { + return Ok(None); + } + apply_runtime_config(runtime, factories, Some(cpex_config(&config.document)?)).await?; + Ok(Some(config.fingerprint)) + } + .await; + if let Err(error) = &result { + set_runtime_failed(runtime, error); + } + result +} + +async fn apply_runtime_config( + runtime: &ArcSwap, + factories: &PluginFactoryRegistry, + config: Option, +) -> Result<(), GatewayPluginRuntimeError> { + let Some(config) = config else { + drop(runtime.swap(Arc::new(RuntimeState::Active(Arc::new(GatewayPluginRuntime::default()))))); + return Ok(()); + }; + drop( + runtime.swap(Arc::new(RuntimeState::Active(Arc::new( + GatewayPluginRuntime::from_config(config, factories).await?, + )))), + ); + Ok(()) +} + +fn set_runtime_failed(runtime: &ArcSwap, error: &GatewayPluginRuntimeError) { + drop(runtime.swap(Arc::new(RuntimeState::Failed(error.to_string())))); +} + +impl CpexRuntimeRegistry { + pub async fn initialize(&self) -> Result>, RuntimeHookError> { + let initial_config = reload_runtime(&self.runtime, self.config_store.as_ref(), &self.factories, None).await?; + Ok(self.start_config_watcher(initial_config)) + } +} + +impl GatewayPluginRuntimeHandle { + pub(crate) fn current(&self) -> Result, ErrorData> { + match self.runtime.load().as_ref() { + RuntimeState::Active(runtime) => Ok(Arc::clone(runtime)), + RuntimeState::Failed(error) => { + tracing::warn!(%error, "rejecting MCP call because CPEX runtime is failed"); + Err(ErrorData { + code: ErrorCode::INTERNAL_ERROR, + message: "Runtime plugin reload failed".into(), + data: None, + }) + }, + } + } +} + +#[cfg(test)] +mod tests; diff --git a/crates/contextforge-data-plane-cpex/src/registry/tests.rs b/crates/contextforge-data-plane-cpex/src/registry/tests.rs new file mode 100644 index 00000000..ee2320bb --- /dev/null +++ b/crates/contextforge-data-plane-cpex/src/registry/tests.rs @@ -0,0 +1,976 @@ +use std::{ + collections::HashMap, + sync::{ + Arc, Mutex, + atomic::{AtomicUsize, Ordering}, + }, + time::Duration, +}; + +use async_trait::async_trait; +use cpex::cpex_core::{ + cmf::{CmfHook, ContentPart, MessagePayload}, + context::PluginContext, + error::{PluginError, PluginViolation}, + factory::{PluginFactory, PluginInstance}, + hooks::{Extensions, HookHandler, PluginResult, TypedHandlerAdapter, types::cmf_hook_names}, + plugin::{Plugin, PluginConfig}, + registry::AnyHookHandler, +}; +use rmcp::model::{ + CallToolRequestParams, CallToolResult, ContentBlock, NumberOrString, ProgressNotificationParam, ProgressToken, + ReadResourceResult, ResourceContents, +}; +use serde_json::{Value, json}; +use tokio::sync::Mutex as TokioMutex; + +use contextforge_data_plane_apis::runtime_plugin_config::{RUNTIME_PLUGIN_CONFIG_VERSION, RuntimePluginConfigDocument}; + +use crate::config::LoadedRuntimePluginConfig; +use crate::{ArgumentsUpdate, CmfPluginFactory, PreHookResult, ToolHookState}; +use rmcp::model::GetPromptRequestParams; + +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); +const TEST_WATCHER_RETRY_COUNT: usize = 20; +const TEST_WATCHER_RETRY_INTERVAL: Duration = Duration::from_millis(20); + +#[derive(Clone, Default)] +struct MemoryConfigStore { + config: Arc>>, + calls: Arc, +} + +impl MemoryConfigStore { + fn with_config(config: RuntimePluginConfigDocument) -> Self { + Self { config: Arc::new(TokioMutex::new(Some(config))), calls: Arc::new(AtomicUsize::new(0)) } + } + + async fn set_config(&self, config: RuntimePluginConfigDocument) { + *self.config.lock().await = Some(config); + } + + async fn clear_config(&self) { + *self.config.lock().await = None; + } + + fn calls(&self) -> usize { + self.calls.load(Ordering::SeqCst) + } +} + +#[async_trait] +impl RuntimePluginConfigStore for MemoryConfigStore { + async fn get_config(&self) -> Result, GatewayPluginRuntimeError> { + self.calls.fetch_add(1, Ordering::SeqCst); + Ok(self.config.lock().await.as_ref().map(loaded_config)) + } +} + +#[derive(Default)] +struct Observations { + pre_calls: usize, + post_calls: usize, + shutdown_calls: usize, + pre_request_id: Option, + post_request_id: Option, +} + +#[derive(Clone, Copy, Default)] +enum PreBehavior { + #[default] + Allow, + Rewrite, + SetContext, +} + +#[derive(Clone, Copy, Default)] +enum PostBehavior { + #[default] + Allow, + Rewrite, + RewriteStreamEvent, + RewriteInvalid, + Deny, + RequireContext, + CountEvents, +} + +struct TestPlugin { + config: PluginConfig, + observations: Arc>, + pre_behavior: PreBehavior, + post_behavior: PostBehavior, +} + +impl TestPlugin { + fn new(name: &str, hooks: Vec<&'static str>) -> Self { + Self { + config: PluginConfig { + name: name.to_owned(), + kind: "test".to_owned(), + hooks: hooks.into_iter().map(str::to_owned).collect(), + ..Default::default() + }, + observations: Arc::new(Mutex::new(Observations::default())), + pre_behavior: PreBehavior::Allow, + post_behavior: PostBehavior::Allow, + } + } + + fn rewrite_from_config(config: PluginConfig) -> Self { + Self { config, ..Self::new("generic-pre", vec![cmf_hook_names::TOOL_PRE_INVOKE]).with_pre_rewrite() } + } + + fn with_pre_rewrite(mut self) -> Self { + self.pre_behavior = PreBehavior::Rewrite; + self + } + + fn with_post_rewrite(mut self) -> Self { + self.post_behavior = PostBehavior::Rewrite; + self + } + + fn with_stream_event_rewrite(mut self) -> Self { + self.post_behavior = PostBehavior::RewriteStreamEvent; + self + } + + fn with_invalid_stream_rewrite(mut self) -> Self { + self.post_behavior = PostBehavior::RewriteInvalid; + self + } + + fn with_post_deny(mut self) -> Self { + self.post_behavior = PostBehavior::Deny; + self + } + + fn with_context_roundtrip(mut self) -> Self { + self.pre_behavior = PreBehavior::SetContext; + self.post_behavior = PostBehavior::RequireContext; + self + } + + fn observations(&self) -> Arc> { + Arc::clone(&self.observations) + } +} + +#[async_trait] +impl Plugin for TestPlugin { + fn config(&self) -> &PluginConfig { + &self.config + } + + async fn shutdown(&self) -> Result<(), Box> { + self.observations.lock().expect("observations lock poisoned").shutdown_calls += 1; + Ok(()) + } +} + +#[allow(clippy::unused_async_trait_impl)] +impl HookHandler for TestPlugin { + async fn handle( + &self, + payload: &MessagePayload, + _extensions: &Extensions, + ctx: &mut PluginContext, + ) -> PluginResult { + let is_post = payload.message.content.iter().any(|part| { + matches!( + part, + ContentPart::ToolResult { .. } | ContentPart::PromptResult { .. } | ContentPart::Resource { .. } + ) + }); + let mut observations = self.observations.lock().expect("observations lock poisoned"); + if is_post { + observations.post_calls += 1; + observations.post_request_id = request_id(payload); + } else { + observations.pre_calls += 1; + observations.pre_request_id = request_id(payload); + } + drop(observations); + + if is_post { + match self.post_behavior { + PostBehavior::Allow => PluginResult::allow(), + PostBehavior::Rewrite => PluginResult::modify_payload(payload.clone()), + PostBehavior::RewriteStreamEvent => { + let mut modified = payload.clone(); + if let Some(ContentPart::ToolResult { content }) = + modified.message.content.iter_mut().find(|part| matches!(part, ContentPart::ToolResult { .. })) + && let Ok(mut progress) = + serde_json::from_value::(content.content.clone()) + { + progress.message = progress.message.map(|message| format!("plugin:{message}")); + content.content = serde_json::to_value(progress).expect("progress serializes"); + } + PluginResult::modify_payload(modified) + }, + PostBehavior::RewriteInvalid => { + let mut modified = payload.clone(); + if let Some(ContentPart::ToolResult { content }) = + modified.message.content.iter_mut().find(|part| matches!(part, ContentPart::ToolResult { .. })) + { + content.content = json!("not-a-stream-event"); + } + PluginResult::modify_payload(modified) + }, + PostBehavior::Deny => PluginResult::deny(PluginViolation::new("post_denied", "post denied")), + PostBehavior::CountEvents => { + let results = payload.message.get_tool_results(); + let result = results.first().expect("tool result"); + let events = ctx.get_global("events").and_then(Value::as_u64).unwrap_or_default(); + if result.content.get("progress").is_some() { + ctx.set_global("events", json!(events + 1)); + PluginResult::allow() + } else if events == 2 { + PluginResult::allow() + } else { + PluginResult::deny(PluginViolation::new("missing_events", "tool event context was lost")) + } + }, + PostBehavior::RequireContext => { + if ctx.get_global("pre_seen") == Some(&json!(true)) { + PluginResult::allow() + } else { + PluginResult::deny( + PluginViolation::new("missing_context", "pre context missing") + .with_proto_error_code(TEST_MISSING_CONTEXT_ERROR_CODE), + ) + } + }, + } + } else { + match self.pre_behavior { + PreBehavior::Allow => PluginResult::allow(), + PreBehavior::Rewrite => { + let mut modified = payload.clone(); + if let Some(ContentPart::ToolCall { content }) = + modified.message.content.iter_mut().find(|part| matches!(part, ContentPart::ToolCall { .. })) + { + content.arguments = HashMap::from([ + ("a".to_owned(), json!(TEST_REWRITTEN_SUM_A)), + ("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 => { + ctx.set_global("pre_seen", json!(true)); + PluginResult::allow() + }, + } + } + } +} + +struct TestPluginFactory { + observations: Arc>, + pre_behavior: PreBehavior, + post_behavior: PostBehavior, +} + +impl TestPluginFactory { + fn from_plugin(plugin: &TestPlugin) -> Self { + Self { + observations: Arc::clone(&plugin.observations), + pre_behavior: plugin.pre_behavior, + post_behavior: plugin.post_behavior, + } + } +} + +impl PluginFactory for TestPluginFactory { + fn create(&self, config: &PluginConfig) -> Result> { + let plugin = Arc::new(TestPlugin { + config: config.clone(), + observations: Arc::clone(&self.observations), + pre_behavior: self.pre_behavior, + post_behavior: self.post_behavior, + }); + let handlers = config + .hooks + .iter() + .filter_map(|hook| { + let hook = crate::factory::supported_cmf_hook_name(hook)?; + Some(( + hook, + Arc::new(TypedHandlerAdapter::::new(Arc::clone(&plugin))) as Arc, + )) + }) + .collect(); + let plugin: Arc = plugin; + Ok(PluginInstance { plugin, handlers }) + } +} + +fn sum_request(a: i64, b: i64) -> CallToolRequestParams { + CallToolRequestParams::new("sum") + .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") +} + +fn config_document(cpex: Value) -> RuntimePluginConfigDocument { + RuntimePluginConfigDocument { + version: RUNTIME_PLUGIN_CONFIG_VERSION, + cpex: serde_json::from_value(cpex).expect("test CPEX config parses"), + } +} + +fn loaded_config(document: &RuntimePluginConfigDocument) -> LoadedRuntimePluginConfig { + LoadedRuntimePluginConfig::decode(serde_json::to_vec(document).expect("test CPEX config serializes")) + .expect("test CPEX config decodes") +} + +fn plugin_config(plugins: &[Arc]) -> RuntimePluginConfigDocument { + config_document(json!({ + "plugins": plugins.iter().map(|plugin| { + json!({ + "name": plugin.config.name.clone(), + "kind": plugin.config.kind.clone(), + "hooks": plugin.config.hooks.clone(), + }) + }).collect::>() + })) +} + +fn expect_runtime_failed(result: Result, ErrorData>) -> ErrorData { + match result { + Ok(_) => panic!("runtime should be failed"), + Err(error) => error, + } +} + +async fn runtime_with_plugin(plugin: &Arc, config: RuntimePluginConfigDocument) -> CpexRuntimeRegistry { + let mut runtime = CpexRuntimeRegistry::with_config_store(Arc::new(MemoryConfigStore::with_config(config))); + runtime.register_factory("test", Box::new(TestPluginFactory::from_plugin(plugin))).expect("test factory registers"); + runtime.initialize().await.expect("runtime initializes"); + runtime +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 1)] +async fn runtime_config_store_is_loaded_on_initialize() { + let config_store = MemoryConfigStore::with_config(config_document(json!({ "plugins": [] }))); + let runtime = CpexRuntimeRegistry::with_config_store(Arc::new(config_store.clone())); + + let handle = runtime.initialize().await.expect("runtime initializes"); + + assert!(handle.is_some()); + assert!(config_store.calls() >= 1); +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 1)] +async fn missing_runtime_plugin_config_is_rejected_on_initialize() { + let runtime = CpexRuntimeRegistry::with_config_store(Arc::new(MemoryConfigStore::default())); + + let error = runtime.initialize().await.expect_err("missing config is rejected"); + + assert_eq!("runtime plugin config is missing", error.to_string()); +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 1)] +async fn invalid_runtime_plugin_config_documents_are_rejected() { + for config in [RuntimePluginConfigDocument { version: 2, cpex: CpexConfig::default() }] { + let runtime = CpexRuntimeRegistry::with_config_store(Arc::new(MemoryConfigStore::with_config(config))); + let error = runtime.initialize().await.expect_err("invalid config is rejected"); + + assert_eq!("runtime plugin config is in wrong format", error.to_string()); + } +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 1)] +async fn unsupported_runtime_plugin_config_is_rejected() { + for cpex in [ + json!({ "plugin_settings": { "routing_enabled": true }, "plugins": [] }), + json!({ "plugin_settings": { "fail_on_plugin_error": true }, "plugins": [] }), + json!({ + "plugins": [{ + "name": "scoped", + "kind": "test", + "hooks": [cmf_hook_names::TOOL_PRE_INVOKE], + "conditions": [{ "tools": ["sum"] }] + }] + }), + json!({ "plugins": [{ "name": "llm", "kind": "test", "hooks": [cmf_hook_names::LLM_INPUT] }] }), + ] { + let runtime = + CpexRuntimeRegistry::with_config_store(Arc::new(MemoryConfigStore::with_config(config_document(cpex)))); + let error = runtime.initialize().await.expect_err("unsupported config is rejected"); + + assert_eq!("runtime plugin config is unsupported", error.to_string()); + } +} + +#[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 resource_pre_hook_runs_for_a_canonical_uri() { + let plugin = Arc::new(TestPlugin::new("resource", vec![cmf_hook_names::RESOURCE_PRE_FETCH])); + let observations = plugin.observations(); + let runtime = runtime_with_plugin(&plugin, plugin_config(&[Arc::clone(&plugin)])).await; + + runtime.handle().before_read_resource("file:///password.env").await.expect("resource pre hook runs"); + + assert_eq!(1, observations.lock().expect("observations lock poisoned").pre_calls); +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 1)] +async fn resource_without_post_hook_keeps_its_decision_across_reload() { + let plugin = Arc::new(TestPlugin::new("resource", vec![cmf_hook_names::RESOURCE_POST_FETCH])); + let observations = plugin.observations(); + let runtime = runtime_with_plugin(&plugin, plugin_config(&[Arc::clone(&plugin)])).await; + runtime.apply_config(None).await.expect("disable hooks"); + let state = runtime.handle().before_read_resource("file:///password.env").await.expect("request starts"); + runtime.apply_config(Some(plugin_config(&[plugin]).cpex)).await.expect("enable hooks"); + let response = ReadResourceResult::new(vec![ResourceContents::text("original", "file:///password.env")]); + state.after_read_resource(response).await.expect("in-flight decision survives reload"); + assert_eq!(0, observations.lock().expect("observations lock poisoned").post_calls); +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 1)] +async fn resource_post_hook_keeps_its_runtime_across_reload() { + let plugin = Arc::new(TestPlugin::new("resource", vec![cmf_hook_names::RESOURCE_POST_FETCH]).with_post_deny()); + let runtime = runtime_with_plugin(&plugin, plugin_config(&[Arc::clone(&plugin)])).await; + let state = runtime.handle().before_read_resource("file:///password.env").await.expect("request starts"); + runtime.apply_config(None).await.expect("disable hooks"); + let response = ReadResourceResult::new(vec![ResourceContents::text("secret", "file:///password.env")]); + let error = state.after_read_resource(response).await.expect_err("captured policy still denies"); + assert_eq!("Plugin denied resource", error.message); +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 1)] +async fn resource_hooks_preserve_context_across_the_backend_call() { + let plugin = Arc::new( + TestPlugin::new("resource", vec![cmf_hook_names::RESOURCE_PRE_FETCH, cmf_hook_names::RESOURCE_POST_FETCH]) + .with_context_roundtrip(), + ); + let observations = plugin.observations(); + let runtime = runtime_with_plugin(&plugin, plugin_config(&[Arc::clone(&plugin)])).await; + let pre = runtime.handle().before_read_resource("file:///password.env").await.expect("resource pre hook runs"); + let response = ReadResourceResult::new(vec![ResourceContents::text("secret", "file:///password.env")]); + + pre.after_read_resource(response).await.expect("resource post hook receives pre context"); + + let observations = observations.lock().expect("observations lock poisoned"); + assert_eq!(1, observations.pre_calls); + assert_eq!(1, observations.post_calls); +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 1)] +async fn runtime_config_loads_registered_factory_plugin() { + let plugin = Arc::new(TestPlugin::new("configured-pre", vec![cmf_hook_names::TOOL_PRE_INVOKE]).with_pre_rewrite()); + let observations = plugin.observations(); + let runtime = runtime_with_plugin(&plugin, plugin_config(&[Arc::clone(&plugin)])).await; + + let result = runtime.before_tool_call(&sum_request(1, 2), "sum", "backend").await.expect("pre hook runs"); + + assert!(matches!(result.arguments, ArgumentsUpdate::Replace(Some(_)))); + assert_eq!(1, observations.lock().expect("observations lock poisoned").pre_calls); +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 1)] +async fn runtime_config_loads_generic_cmf_factory_plugin() { + let config = config_document(json!({ + "plugins": [{ + "name": "generic-pre", + "kind": "generic", + "hooks": [cmf_hook_names::TOOL_PRE_INVOKE] + }] + })); + 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.before_tool_call(&sum_request(1, 2), "sum", "backend").await.expect("pre hook runs"); + + assert!(matches!(result.arguments, ArgumentsUpdate::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, ArgumentsUpdate::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, ArgumentsUpdate::Replace(Some(_)))); + assert!(matches!(prompt.arguments, ArgumentsUpdate::Replace(Some(_)))); +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 1)] +async fn runtime_reload_replaces_and_clears_current_runtime() { + let plugin = Arc::new(TestPlugin::new("configured-pre", vec![cmf_hook_names::TOOL_PRE_INVOKE]).with_pre_rewrite()); + let observations = plugin.observations(); + let config_store = MemoryConfigStore::with_config(config_document(json!({ "plugins": [] }))); + let mut runtime = CpexRuntimeRegistry::with_config_store(Arc::new(config_store.clone())); + runtime + .register_factory("test", Box::new(TestPluginFactory::from_plugin(&plugin))) + .expect("test factory registers"); + runtime.initialize().await.expect("runtime initializes"); + + let result = runtime.before_tool_call(&sum_request(1, 2), "sum", "backend").await.expect("pre hook skips"); + assert!(matches!(result.arguments, ArgumentsUpdate::Unchanged)); + + config_store.set_config(plugin_config(&[Arc::clone(&plugin)])).await; + runtime.reload().await.expect("runtime reloads"); + let result = runtime.before_tool_call(&sum_request(1, 2), "sum", "backend").await.expect("pre hook runs"); + assert!(matches!(result.arguments, ArgumentsUpdate::Replace(Some(_)))); + + config_store.set_config(config_document(json!({ "plugins": [] }))).await; + runtime.reload().await.expect("runtime reloads"); + let result = runtime.before_tool_call(&sum_request(1, 2), "sum", "backend").await.expect("pre hook skips"); + assert!(matches!(result.arguments, ArgumentsUpdate::Unchanged)); + assert_eq!(1, observations.lock().expect("observations lock poisoned").pre_calls); +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 1)] +async fn failed_runtime_reload_rejects_new_calls_until_valid_reload() { + let plugin = Arc::new(TestPlugin::new("configured-pre", vec![cmf_hook_names::TOOL_PRE_INVOKE]).with_pre_rewrite()); + let observations = plugin.observations(); + let config_store = MemoryConfigStore::with_config(plugin_config(&[Arc::clone(&plugin)])); + let mut runtime = CpexRuntimeRegistry::with_config_store(Arc::new(config_store.clone())); + runtime + .register_factory("test", Box::new(TestPluginFactory::from_plugin(&plugin))) + .expect("test factory registers"); + runtime.initialize().await.expect("runtime initializes"); + + config_store.set_config(RuntimePluginConfigDocument { version: 2, cpex: CpexConfig::default() }).await; + runtime.reload().await.expect_err("invalid reload fails"); + let error = expect_runtime_failed(runtime.before_tool_call(&sum_request(1, 2), "sum", "backend").await); + assert_eq!(ErrorCode::INTERNAL_ERROR, error.code); + assert_eq!("Runtime plugin reload failed", error.message); + + config_store.clear_config().await; + runtime.reload().await.expect_err("missing reload fails"); + let error = expect_runtime_failed(runtime.before_tool_call(&sum_request(1, 2), "sum", "backend").await); + assert_eq!(ErrorCode::INTERNAL_ERROR, error.code); + assert_eq!("Runtime plugin reload failed", error.message); + + config_store.set_config(plugin_config(&[Arc::clone(&plugin)])).await; + runtime.reload().await.expect("runtime recovers"); + let result = runtime.before_tool_call(&sum_request(1, 2), "sum", "backend").await.expect("pre hook runs"); + assert!(matches!(result.arguments, ArgumentsUpdate::Replace(Some(_)))); + assert_eq!(1, observations.lock().expect("observations lock poisoned").pre_calls); +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 1)] +async fn combined_plugin_preserves_context_from_pre_to_post_across_replacement() { + let plugin = Arc::new( + TestPlugin::new("context", vec![cmf_hook_names::TOOL_PRE_INVOKE, cmf_hook_names::TOOL_POST_INVOKE]) + .with_context_roundtrip(), + ); + let observations = plugin.observations(); + let config_store = MemoryConfigStore::with_config(plugin_config(&[Arc::clone(&plugin)])); + let mut runtime = CpexRuntimeRegistry::with_config_store(Arc::new(config_store.clone())); + runtime + .register_factory("test", Box::new(TestPluginFactory::from_plugin(&plugin))) + .expect("test factory registers"); + runtime.initialize().await.expect("runtime initializes"); + + let pre = runtime.before_tool_call(&sum_request(1, 2), "sum", "backend").await.expect("pre hook runs"); + config_store.set_config(config_document(json!({ "plugins": [] }))).await; + runtime.reload().await.expect("runtime reloads"); + let response = CallToolResult::success(vec![ContentBlock::text("3")]); + runtime.after_tool_call("sum", response, pre.state).await.expect("post hook runs"); + + let observations = observations.lock().expect("observations lock poisoned"); + assert_eq!(1, observations.post_calls); + assert_eq!(observations.pre_request_id, observations.post_request_id); +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 1)] +async fn post_only_runtime_does_not_apply_new_post_hook_to_in_flight_call() { + let plugin = Arc::new(TestPlugin::new("post", vec![cmf_hook_names::TOOL_POST_INVOKE]).with_post_rewrite()); + let observations = plugin.observations(); + let config_store = MemoryConfigStore::with_config(config_document(json!({ "plugins": [] }))); + let mut runtime = CpexRuntimeRegistry::with_config_store(Arc::new(config_store.clone())); + runtime + .register_factory("test", Box::new(TestPluginFactory::from_plugin(&plugin))) + .expect("test factory registers"); + runtime.initialize().await.expect("runtime initializes"); + + let pre = runtime.before_tool_call(&sum_request(1, 2), "sum", "backend").await.expect("pre hook skips"); + config_store.set_config(plugin_config(&[Arc::clone(&plugin)])).await; + runtime.reload().await.expect("runtime reloads"); + let response = CallToolResult::success(vec![ContentBlock::text("3")]); + runtime.after_tool_call("sum", response, pre.state).await.expect("post hook skips"); + + assert_eq!(0, observations.lock().expect("observations lock poisoned").post_calls); +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 1)] +async fn disabled_runtime_does_not_create_tool_post_state() { + let plugin = Arc::new(TestPlugin::new("post", vec![cmf_hook_names::TOOL_POST_INVOKE]).with_stream_event_rewrite()); + let observations = plugin.observations(); + let runtime = runtime_with_plugin(&plugin, plugin_config(&[Arc::clone(&plugin)])).await; + + runtime.apply_config(None).await.expect("disable hooks"); + let pre = runtime.before_tool_call(&sum_request(1, 2), "sum", "backend").await.expect("request starts"); + + assert!(pre.state.is_none()); + assert_eq!(0, observations.lock().expect("observations lock poisoned").post_calls); +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 1)] +async fn stream_event_is_rewritten_by_post_hook() { + let plugin = Arc::new(TestPlugin::new("post", vec![cmf_hook_names::TOOL_POST_INVOKE]).with_stream_event_rewrite()); + let observations = plugin.observations(); + let runtime = runtime_with_plugin(&plugin, plugin_config(&[Arc::clone(&plugin)])).await; + + let pre = runtime.before_tool_call(&sum_request(1, 2), "sum", "backend").await.expect("pre state is created"); + let event = pre.state.expect("post state").after_stream_event(progress_event()).await.expect("event passes"); + + assert_eq!(Some("plugin:step 1/2"), event.expect("event is kept").message.as_deref()); + assert_eq!(1, observations.lock().expect("observations lock poisoned").post_calls); +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 1)] +async fn denied_stream_event_is_dropped() { + let plugin = Arc::new(TestPlugin::new("post", vec![cmf_hook_names::TOOL_POST_INVOKE]).with_post_deny()); + let runtime = runtime_with_plugin(&plugin, plugin_config(&[Arc::clone(&plugin)])).await; + + let pre = runtime.before_tool_call(&sum_request(1, 2), "sum", "backend").await.expect("pre state is created"); + let event = + pre.state.expect("post state").after_stream_event(progress_event()).await.expect("deny drops the event"); + + assert!(event.is_none()); +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 1)] +async fn invalid_stream_event_rewrite_is_rejected() { + let plugin = + Arc::new(TestPlugin::new("post", vec![cmf_hook_names::TOOL_POST_INVOKE]).with_invalid_stream_rewrite()); + let runtime = runtime_with_plugin(&plugin, plugin_config(&[Arc::clone(&plugin)])).await; + + let pre = runtime.before_tool_call(&sum_request(1, 2), "sum", "backend").await.expect("pre state is created"); + let error = pre + .state + .expect("post state") + .after_stream_event(progress_event()) + .await + .expect_err("invalid rewrite is rejected"); + + assert_eq!(ErrorCode::INVALID_PARAMS, error.code); +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 1)] +async fn replaced_runtime_shutdowns_on_drop() { + let plugin = Arc::new(TestPlugin::new("pre", vec![cmf_hook_names::TOOL_PRE_INVOKE]).with_pre_rewrite()); + let observations = plugin.observations(); + let config_store = MemoryConfigStore::with_config(plugin_config(&[Arc::clone(&plugin)])); + let mut runtime = CpexRuntimeRegistry::with_config_store(Arc::new(config_store.clone())); + runtime + .register_factory("test", Box::new(TestPluginFactory::from_plugin(&plugin))) + .expect("test factory registers"); + runtime.initialize().await.expect("runtime initializes"); + + config_store.set_config(config_document(json!({ "plugins": [] }))).await; + runtime.reload().await.expect("runtime reloads"); + + for _ in 0..TEST_SHUTDOWN_RETRY_COUNT { + if observations.lock().expect("observations lock poisoned").shutdown_calls > 0 { + return; + } + tokio::time::sleep(TEST_SHUTDOWN_RETRY_INTERVAL).await; + } + panic!("replaced runtime did not shut down"); +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 1)] +async fn watcher_applies_config_changes() { + let plugin = Arc::new(TestPlugin::new("configured-pre", vec![cmf_hook_names::TOOL_PRE_INVOKE]).with_pre_rewrite()); + let observations = plugin.observations(); + let config_store = MemoryConfigStore::with_config(config_document(json!({ "plugins": [] }))); + let mut runtime = + CpexRuntimeRegistry::with_config_store_interval(Arc::new(config_store.clone()), TEST_WATCHER_INTERVAL); + runtime + .register_factory("test", Box::new(TestPluginFactory::from_plugin(&plugin))) + .expect("test factory registers"); + let handle = runtime.initialize().await.expect("runtime initializes"); + assert!(handle.is_some()); + + config_store.set_config(plugin_config(&[Arc::clone(&plugin)])).await; + for _ in 0..TEST_WATCHER_RETRY_COUNT { + let result = runtime.before_tool_call(&sum_request(1, 2), "sum", "backend").await.expect("pre hook runs"); + if matches!(result.arguments, ArgumentsUpdate::Replace(Some(_))) { + config_store.clear_config().await; + tokio::time::sleep(TEST_WATCHER_INTERVAL + TEST_WATCHER_RETRY_INTERVAL).await; + let error = expect_runtime_failed(runtime.before_tool_call(&sum_request(1, 2), "sum", "backend").await); + assert_eq!(ErrorCode::INTERNAL_ERROR, error.code); + + config_store.set_config(plugin_config(&[Arc::clone(&plugin)])).await; + for _ in 0..TEST_WATCHER_RETRY_COUNT { + if let Ok(result) = runtime.before_tool_call(&sum_request(1, 2), "sum", "backend").await + && matches!(result.arguments, ArgumentsUpdate::Replace(Some(_))) + { + assert_eq!(2, observations.lock().expect("observations lock poisoned").pre_calls); + return; + } + tokio::time::sleep(TEST_WATCHER_RETRY_INTERVAL).await; + } + panic!("config watcher did not recover from missing plugin config"); + } + tokio::time::sleep(TEST_WATCHER_RETRY_INTERVAL).await; + } + panic!("config watcher did not apply plugin config"); +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 1)] +async fn initialize_without_config_store_returns_no_watcher() { + let runtime = CpexRuntimeRegistry::default(); + + let handle = runtime.initialize().await.expect("runtime initializes"); + + assert!(handle.is_none()); +} + +#[cfg(test)] +impl CpexRuntimeRegistry { + fn with_config_store(config_store: Arc) -> Self { + Self { config_store: Some(config_store), ..Self::default() } + } + + fn with_config_store_interval(config_store: Arc, watcher_interval: Duration) -> Self { + Self { config_store: Some(config_store), watcher_interval, ..Self::default() } + } + + async fn before_tool_call( + &self, + request: &CallToolRequestParams, + tool_name: &str, + backend_name: &str, + ) -> Result, ErrorData> { + self.handle().before_tool_call(request, tool_name, backend_name).await + } + + async fn after_tool_call( + &self, + _tool_name: &str, + response: CallToolResult, + state: Option, + ) -> Result { + match state { + Some(state) => state.after_tool_call(response).await, + None => Ok(response), + } + } +} + +fn request_id(payload: &MessagePayload) -> Option { + payload.message.content.iter().find_map(|part| match part { + ContentPart::ToolCall { content } => Some(content.tool_call_id.clone()), + ContentPart::ToolResult { content } => Some(content.tool_call_id.clone()), + ContentPart::PromptRequest { content } => Some(content.prompt_request_id.clone()), + ContentPart::PromptResult { content } => Some(content.prompt_request_id.clone()), + ContentPart::ResourceRef { content } => Some(content.resource_request_id.clone()), + ContentPart::Resource { content } => Some(content.resource_request_id.clone()), + _ => None, + }) +} + +#[tokio::test] +async fn hook_combinations_preserve_correlation_for_each_operation() { + use crate::cmf::Operation; + use rmcp::model::{GetPromptResult, PromptMessage, Role}; + + for operation in Operation::ALL { + for (pre_enabled, post_enabled) in [(false, false), (true, false), (false, true), (true, true)] { + let [pre, post] = operation.hooks(); + let hooks = [(pre, pre_enabled), (post, post_enabled)] + .into_iter() + .filter_map(|(hook, enabled)| enabled.then_some(hook)) + .collect(); + let plugin = Arc::new(TestPlugin::new("combinations", hooks)); + let observations = plugin.observations(); + let runtime = runtime_with_plugin(&plugin, plugin_config(&[Arc::clone(&plugin)])).await; + let handle = runtime.handle(); + + match operation { + Operation::Tool => { + let pre = handle.before_tool_call(&sum_request(1, 2), "sum", "backend").await.expect("tool starts"); + assert_eq!(post_enabled, pre.state.is_some()); + if let Some(state) = pre.state { + state.after_tool_call(CallToolResult::success(vec![])).await.expect("tool finishes"); + } + }, + Operation::Prompt => { + let pre = handle + .before_get_prompt(&review_request("weather"), "review", "backend") + .await + .expect("prompt starts"); + assert_eq!(post_enabled, pre.state.is_some()); + if let Some(state) = pre.state { + state + .after_get_prompt(GetPromptResult::new(vec![PromptMessage::new_text( + Role::User, + "weather", + )])) + .await + .expect("prompt finishes"); + } + }, + Operation::Resource => { + let state = handle.before_read_resource("file:///test").await.expect("resource starts"); + state + .after_read_resource(ReadResourceResult::new(vec![ResourceContents::text( + "weather", + "file:///test", + )])) + .await + .expect("resource finishes"); + }, + } + let observations = observations.lock().expect("observations lock"); + assert_eq!(usize::from(pre_enabled), observations.pre_calls); + assert_eq!(usize::from(post_enabled), observations.post_calls); + if pre_enabled && post_enabled { + assert!(observations.pre_request_id.is_some()); + assert_eq!(observations.pre_request_id, observations.post_request_id); + } + } + } +} + +#[tokio::test] +async fn prompt_context_and_policy_survive_a_failed_reload() { + use rmcp::model::{GetPromptResult, PromptMessage, Role}; + + let plugin = Arc::new( + TestPlugin::new("prompt-context", vec![cmf_hook_names::PROMPT_PRE_FETCH, cmf_hook_names::PROMPT_POST_FETCH]) + .with_context_roundtrip(), + ); + let config_store = MemoryConfigStore::with_config(plugin_config(&[Arc::clone(&plugin)])); + let mut runtime = CpexRuntimeRegistry::with_config_store(Arc::new(config_store.clone())); + runtime.register_factory("test", Box::new(TestPluginFactory::from_plugin(&plugin))).expect("factory registers"); + runtime.initialize().await.expect("runtime initializes"); + let handle = runtime.handle(); + let request = review_request("weather"); + let pre = handle.before_get_prompt(&request, "review", "backend").await.expect("prompt starts"); + config_store.clear_config().await; + runtime.reload().await.expect_err("missing config fails reload"); + assert!(handle.before_get_prompt(&request, "review", "backend").await.is_err()); + assert!(handle.before_read_resource("file:///test").await.is_err()); + pre.state + .expect("prompt state") + .after_get_prompt(GetPromptResult::new(vec![PromptMessage::new_text(Role::User, "weather")])) + .await + .expect("original prompt context remains usable"); + let observations = plugin.observations.lock().expect("observations lock"); + assert_eq!(1, observations.post_calls); + assert_eq!(observations.pre_request_id, observations.post_request_id); +} + +#[tokio::test] +async fn concurrent_tool_events_share_context_with_the_final_response_after_reload() { + let mut plugin = TestPlugin::new("event-context", vec![cmf_hook_names::TOOL_POST_INVOKE]); + plugin.post_behavior = PostBehavior::CountEvents; + let plugin = Arc::new(plugin); + let runtime = runtime_with_plugin(&plugin, plugin_config(&[Arc::clone(&plugin)])).await; + let pre = runtime.before_tool_call(&sum_request(1, 2), "sum", "backend").await.expect("tool starts"); + let state = pre.state.expect("tool state"); + runtime.apply_config(None).await.expect("disable hooks for new calls"); + let (first, second) = + tokio::join!(state.after_stream_event(progress_event()), state.after_stream_event(progress_event())); + assert!(first.expect("first event").is_some()); + assert!(second.expect("second event").is_some()); + state.after_tool_call(CallToolResult::success(vec![])).await.expect("both event updates reach the final hook"); + assert_eq!(3, plugin.observations.lock().expect("observations lock").post_calls); +} + +#[tokio::test] +async fn dropping_the_last_in_flight_state_releases_the_replaced_runtime() { + let plugin = Arc::new(TestPlugin::new("pending", vec![cmf_hook_names::TOOL_POST_INVOKE])); + let runtime = runtime_with_plugin(&plugin, plugin_config(&[Arc::clone(&plugin)])).await; + let pre = runtime.before_tool_call(&sum_request(1, 2), "sum", "backend").await.expect("tool starts"); + let state = pre.state.expect("tool state"); + let event_state = state.clone(); + runtime.apply_config(None).await.expect("replace runtime"); + drop(state); + tokio::task::yield_now().await; + assert_eq!(0, plugin.observations.lock().expect("observations lock").shutdown_calls); + drop(event_state); + for _ in 0..TEST_SHUTDOWN_RETRY_COUNT { + if plugin.observations.lock().expect("observations lock").shutdown_calls == 1 { + return; + } + tokio::time::sleep(TEST_SHUTDOWN_RETRY_INTERVAL).await; + } + panic!("abandoned request retained its runtime"); +} diff --git a/crates/contextforge-data-plane-cpex/src/resources/mod.rs b/crates/contextforge-data-plane-cpex/src/resources/mod.rs new file mode 100644 index 00000000..190a16fc --- /dev/null +++ b/crates/contextforge-data-plane-cpex/src/resources/mod.rs @@ -0,0 +1,161 @@ +use base64::{Engine as _, prelude::BASE64_STANDARD}; +use cpex::cpex_core::cmf::{ + ContentPart, MessagePayload, Resource as CmfResource, ResourceReference, ResourceType, Role, +}; +use rmcp::{ + ErrorData, + model::{ReadResourceResult, ResourceContents}, +}; + +use crate::{ + GatewayPluginRuntimeHandle, + cmf::{CmfResponse, Operation, message_payload}, + runtime::CallState, +}; + +fn resource_request_payload(resource_uri: &str, resource_request_id: &str) -> MessagePayload { + message_payload( + Role::User, + vec![ContentPart::ResourceRef { + content: ResourceReference { + resource_request_id: resource_request_id.to_owned(), + uri: resource_uri.to_owned(), + name: None, + resource_type: ResourceType::Uri, + range_start: None, + range_end: None, + selector: None, + }, + }], + ) +} + +fn resource_result_payload(response: &ReadResourceResult, resource_request_id: &str) -> Option { + let content = response + .contents + .iter() + .map(|content| { + cmf_resource_content(content, resource_request_id).map(|content| ContentPart::Resource { content }) + }) + .collect::>>()?; + Some(message_payload(Role::Assistant, content)) +} + +fn cmf_resource_content(content: &ResourceContents, resource_request_id: &str) -> Option { + let (uri, mime_type, text, blob) = match content { + ResourceContents::TextResourceContents { uri, mime_type, text, .. } => { + (uri.clone(), mime_type.clone(), Some(text.clone()), None) + }, + ResourceContents::BlobResourceContents { uri, mime_type, blob, .. } => { + (uri.clone(), mime_type.clone(), None, Some(BASE64_STANDARD.decode(blob).ok()?)) + }, + _ => return None, + }; + Some(CmfResource { + resource_request_id: resource_request_id.to_owned(), + uri, + resource_type: ResourceType::Uri, + content: text, + blob, + mime_type, + ..Default::default() + }) +} + +fn resource_result_response(mut original: ReadResourceResult, payload: &MessagePayload) -> Option { + // Resource post hooks replace each resource's content, not the read envelope. + if payload.message.content.len() != original.contents.len() { + return None; + } + for (original, modified) in original.contents.iter_mut().zip(&payload.message.content) { + let ContentPart::Resource { content } = modified else { return None }; + let meta = match original { + ResourceContents::TextResourceContents { meta, .. } + | ResourceContents::BlobResourceContents { meta, .. } => meta.clone(), + _ => return None, + }; + *original = match (&content.content, &content.blob) { + (Some(text), _) => ResourceContents::TextResourceContents { + uri: content.uri.clone(), + mime_type: content.mime_type.clone(), + text: text.clone(), + meta, + }, + (None, Some(bytes)) => { + let blob = match original { + ResourceContents::BlobResourceContents { blob, .. } + if BASE64_STANDARD.decode(blob.as_bytes()).ok().as_ref() == Some(bytes) => + { + blob.clone() + }, + _ => BASE64_STANDARD.encode(bytes), + }; + ResourceContents::BlobResourceContents { + uri: content.uri.clone(), + mime_type: content.mime_type.clone(), + blob, + meta, + } + }, + _ => return None, + }; + } + Some(original) +} + +/// Captures the resource URI edit and post-hook decision for one request. +pub struct ResourceHookState { + rewritten_uri: Option, + call: Option, +} + +impl ResourceHookState { + pub fn rewritten_uri(&self) -> Option<&str> { + self.rewritten_uri.as_deref() + } + + pub async fn after_read_resource(self, response: ReadResourceResult) -> Result { + match self.call { + Some(mut call) => call.after(response).await, + None => Ok(response), + } + } +} + +impl GatewayPluginRuntimeHandle { + pub async fn before_read_resource(&self, resource_uri: &str) -> Result { + let (rewritten_uri, call) = self + .current()? + .before( + Operation::Resource, + "", + |id| resource_request_payload(resource_uri, id), + |payload, _| { + let [ContentPart::ResourceRef { content }] = payload.message.content.as_slice() else { + return Err(ErrorData::internal_error("Plugin returned an invalid resource request", None)); + }; + Ok(Some(content.uri.clone())) + }, + ) + .await?; + Ok(ResourceHookState { rewritten_uri, call }) + } +} + +impl CmfResponse for ReadResourceResult { + const OPERATION: Operation = Operation::Resource; + + fn to_payload(&self, _name: &str, id: &str) -> Result { + resource_result_payload(self, id) + .ok_or_else(|| ErrorData::internal_error("Resource response contains an unsupported content type", None)) + } + + fn apply_payload(self, payload: &MessagePayload, _name: &str, _id: &str) -> Result { + resource_result_response(self, payload).ok_or_else(|| { + ErrorData::internal_error("Plugin returned a resource result the gateway cannot apply", None) + }) + } +} + +#[cfg(test)] +mod tests; diff --git a/crates/contextforge-data-plane-cpex/src/resources/tests.rs b/crates/contextforge-data-plane-cpex/src/resources/tests.rs new file mode 100644 index 00000000..24083776 --- /dev/null +++ b/crates/contextforge-data-plane-cpex/src/resources/tests.rs @@ -0,0 +1,142 @@ +use super::*; + +#[test] +fn resource_result_response_applies_text_changes() { + let original = + ReadResourceResult::new(vec![ResourceContents::text("AWS_ACCESS_KEY_ID=secret", "file:///password.env")]); + let mut payload = resource_result_payload(&original, "resource-1").expect("resource result is supported"); + let ContentPart::Resource { content } = &mut payload.message.content[0] else { + panic!("expected resource content"); + }; + content.content = Some("AWS_ACCESS_KEY_ID=[redacted]".to_owned()); + + let result = resource_result_response(original, &payload).expect("text edit applies"); + + let ResourceContents::TextResourceContents { text, uri, .. } = &result.contents[0] else { + panic!("expected text resource"); + }; + assert_eq!("AWS_ACCESS_KEY_ID=[redacted]", text); + assert_eq!("file:///password.env", uri); +} + +#[test] +fn resource_result_response_decodes_and_applies_blob_changes() { + let wire_blob = BASE64_STANDARD.encode(b"AWS_ACCESS_KEY_ID=secret"); + let original = ReadResourceResult::new(vec![ + ResourceContents::blob(wire_blob.clone(), "file:///password.bin").with_mime_type("application/octet-stream"), + ]); + let mut payload = resource_result_payload(&original, "resource-1").expect("valid blob is supported"); + let ContentPart::Resource { content } = &mut payload.message.content[0] else { + panic!("expected resource content"); + }; + assert_eq!(Some(b"AWS_ACCESS_KEY_ID=secret".as_slice()), content.blob.as_deref()); + content.blob = Some(b"AWS_ACCESS_KEY_ID=[redacted]".to_vec()); + + let result = resource_result_response(original, &payload).expect("blob edit applies"); + + let ResourceContents::BlobResourceContents { blob, uri, .. } = &result.contents[0] else { + panic!("expected blob resource"); + }; + assert_eq!(b"AWS_ACCESS_KEY_ID=[redacted]", BASE64_STANDARD.decode(blob).expect("valid base64").as_slice()); + assert_eq!("file:///password.bin", uri); + assert_ne!(&wire_blob, blob); +} + +#[test] +fn resource_result_response_preserves_unchanged_blob_wire_value() { + let wire_blob = BASE64_STANDARD.encode(b"unchanged"); + let original = ReadResourceResult::new(vec![ResourceContents::blob(&wire_blob, "file:///image.bin")]); + let payload = resource_result_payload(&original, "resource-1").expect("valid blob is supported"); + + let result = resource_result_response(original, &payload).expect("unchanged blob applies"); + + let ResourceContents::BlobResourceContents { blob, .. } = &result.contents[0] else { + panic!("expected blob resource"); + }; + assert_eq!(&wire_blob, blob); +} + +#[test] +fn resource_result_payload_rejects_invalid_base64_blob() { + let original = ReadResourceResult::new(vec![ResourceContents::blob("not base64!", "file:///image.bin")]); + + assert!(resource_result_payload(&original, "resource-1").is_none()); +} + +#[test] +fn resource_result_allows_mime_uri_and_content_type_changes() { + let original: ReadResourceResult = serde_json::from_value(serde_json::json!({ + "_meta": {"response": "preserved"}, + "contents": [ + {"uri": "file:///a", "mimeType": "text/plain", "text": "original", "_meta": {"item": 1}}, + {"uri": "file:///b", "mimeType": "application/octet-stream", "blob": "YmluYXJ5", "_meta": {"item": 2}} + ] + })) + .expect("valid resource response"); + let mut payload = resource_result_payload(&original, "resource-1").expect("resource payload"); + for (index, part) in payload.message.content.iter_mut().enumerate() { + let ContentPart::Resource { content } = part else { panic!("resource content") }; + content.uri = format!("file:///changed-{index}"); + if index == 0 { + content.content = None; + content.blob = Some(b"binary edit".to_vec()); + content.mime_type = Some("application/octet-stream".to_owned()); + } else { + content.blob = None; + content.content = Some("text edit".to_owned()); + content.mime_type = Some("text/plain".to_owned()); + } + } + let actual = serde_json::to_value(resource_result_response(original, &payload).expect("valid changes apply")) + .expect("response serializes"); + assert_eq!(serde_json::json!({"response": "preserved"}), actual["_meta"]); + assert_eq!("file:///changed-0", actual["contents"][0]["uri"]); + assert_eq!("application/octet-stream", actual["contents"][0]["mimeType"]); + assert_eq!(BASE64_STANDARD.encode(b"binary edit"), actual["contents"][0]["blob"]); + assert_eq!(1, actual["contents"][0]["_meta"]["item"]); + assert_eq!("file:///changed-1", actual["contents"][1]["uri"]); + assert_eq!("text/plain", actual["contents"][1]["mimeType"]); + assert_eq!("text edit", actual["contents"][1]["text"]); + assert_eq!(2, actual["contents"][1]["_meta"]["item"]); +} + +#[test] +fn resource_result_ignores_cmf_fields_that_are_not_mcp_content() { + let original = ReadResourceResult::new(vec![ResourceContents::text("original", "file:///a")]); + let mut payload = resource_result_payload(&original, "resource-1").expect("resource payload"); + payload.message.schema_version = "plugin value".to_owned(); + payload.message.role = Role::User; + payload.message.channel = Some(cpex::cpex_core::cmf::Channel::Analysis); + let ContentPart::Resource { content } = &mut payload.message.content[0] else { panic!("resource content") }; + content.resource_request_id = "plugin value".to_owned(); + content.name = Some("display name".to_owned()); + content.description = Some("description".to_owned()); + content.size_bytes = Some(8); + content.version = Some("v2".to_owned()); + content.annotations.insert("note".to_owned(), serde_json::json!("annotation")); + content.content = Some("redacted".to_owned()); + let result = resource_result_response(original, &payload).expect("MCP content remains usable"); + let ResourceContents::TextResourceContents { text, .. } = &result.contents[0] else { panic!("text content") }; + assert_eq!("redacted", text); +} + +#[test] +fn resource_result_prefers_text_when_both_content_fields_are_present() { + let original = ReadResourceResult::new(vec![ResourceContents::blob("YmluYXJ5", "file:///a")]); + let mut payload = resource_result_payload(&original, "resource-1").expect("resource payload"); + let ContentPart::Resource { content } = &mut payload.message.content[0] else { panic!("resource content") }; + content.content = Some("text replacement".to_owned()); + let result = + resource_result_response(original, &payload).expect("text takes precedence, as in the built-in serializer"); + let ResourceContents::TextResourceContents { text, .. } = &result.contents[0] else { panic!("text resource") }; + assert_eq!("text replacement", text); +} + +#[test] +fn resource_result_rejects_content_without_a_valid_mcp_representation() { + let original = ReadResourceResult::new(vec![ResourceContents::text("original", "file:///a")]); + let mut payload = resource_result_payload(&original, "resource-1").expect("resource payload"); + let ContentPart::Resource { content } = &mut payload.message.content[0] else { panic!("resource content") }; + content.content = None; + assert!(resource_result_response(original, &payload).is_none()); +} diff --git a/crates/contextforge-data-plane-cpex/src/runtime.rs b/crates/contextforge-data-plane-cpex/src/runtime.rs index 2d4b4914..d82792b6 100644 --- a/crates/contextforge-data-plane-cpex/src/runtime.rs +++ b/crates/contextforge-data-plane-cpex/src/runtime.rs @@ -1,3 +1,4 @@ +//! CPEX manager lifecycle and the common pre/post execution path. use std::sync::{ Arc, atomic::{AtomicU64, Ordering}, @@ -9,29 +10,15 @@ use cpex::cpex_core::{ context::PluginContextTable, executor::PipelineResult, factory::PluginFactoryRegistry, - hooks::{payload::Extensions, types::cmf_hook_names}, + hooks::payload::Extensions, manager::PluginManager, }; -use rmcp::{ - ErrorData, - model::{CallToolRequestParams, CallToolResult, GetPromptRequestParams, GetPromptResult, ReadResourceResult}, - serde::{Serialize, de::DeserializeOwned}, -}; -use tokio::sync::Mutex; +use rmcp::ErrorData; use crate::{ - cmf::{ - prompt_request_payload, prompt_result_payload, resource_request_payload, resource_result_payload, - tool_call_payload, tool_json_result_payload, tool_result_payload, - }, + cmf::{CmfResponse, Operation, modified_message_payload, plugin_denied_error}, error::GatewayPluginRuntimeError, factory::supported_cmf_hook_name, - hooks::{PromptPreFetchResult, RuntimeHookState, ToolArgumentsUpdate, ToolPreCallResult}, - pipeline::{ - effective_post_json, effective_post_prompt_result, effective_post_resource_result, effective_post_result, - effective_pre_args, effective_pre_prompt_args, effective_pre_resource_uri, log_pipeline_errors, - plugin_denied_error, - }, }; #[derive(Default)] @@ -40,95 +27,117 @@ struct HookPair { post: bool, } -#[derive(Default)] -struct HookPresence { - tool: HookPair, - prompt: HookPair, - resource: HookPair, -} - #[derive(Default)] pub(crate) struct GatewayPluginRuntime { manager: PluginManager, - hooks: HookPresence, + hooks: [HookPair; 3], } -struct ToolCallState { +/// Pins the selected runtime and correlation context until the request finishes. +/// Only tool calls wrap this in a mutex, because events also update their context. +pub(crate) struct CallState { + runtime: Arc, context_table: PluginContextTable, - tool_call_id: String, + name: String, + id: String, } -type SharedToolCallState = Mutex; - static CORRELATION_ID: AtomicU64 = AtomicU64::new(1); -fn next_tool_call_id() -> String { - format!("gateway-tool-call-{}", CORRELATION_ID.fetch_add(1, Ordering::Relaxed)) -} - -fn new_tool_call_state() -> RuntimeHookState { - Arc::new(Mutex::new(ToolCallState { - context_table: PluginContextTable::default(), - tool_call_id: next_tool_call_id(), - })) -} - -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 }) -} - -pub(crate) struct ResourceCallState { - context_table: PluginContextTable, - resource_request_id: String, -} - -fn next_resource_request_id() -> String { - format!("gateway-resource-request-{}", CORRELATION_ID.fetch_add(1, Ordering::Relaxed)) -} - impl GatewayPluginRuntime { - pub(crate) fn has_post_hook(&self) -> bool { - self.hooks.tool.post - } - - pub(crate) fn has_prompt_post_hook(&self) -> bool { - self.hooks.prompt.post - } - pub(crate) async fn from_config( config: CpexConfig, factories: &PluginFactoryRegistry, ) -> Result { validate_gateway_supported_config(&config)?; - 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), - }, - resource: HookPair { - pre: declares(&config, cmf_hook_names::RESOURCE_PRE_FETCH), - post: declares(&config, cmf_hook_names::RESOURCE_POST_FETCH), - }, - }; + let hooks = Operation::ALL.map(|operation| { + let [pre, post] = operation.hooks().map(|name| declares(&config, name)); + HookPair { pre, post } + }); 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, hooks }) } + + pub(crate) async fn before( + self: &Arc, + operation: Operation, + name: &str, + payload: impl FnOnce(&str) -> MessagePayload, + update: impl FnOnce(&MessagePayload, &str) -> Result, + ) -> Result<(U, Option), ErrorData> { + let hooks = &self.hooks[operation as usize]; + if !hooks.pre && !hooks.post { + return Ok((U::default(), None)); + } + + let id = format!("{}-{}", operation.id_prefix(), CORRELATION_ID.fetch_add(1, Ordering::Relaxed)); + let (update, context_table) = if hooks.pre { + let result = self.invoke(operation.hooks()[0], payload(&id), None).await; + if result.is_denied() { + return Err(plugin_denied_error(operation.subject(), result)); + } + let update = match modified_message_payload(&result) { + Some(payload) => update(payload, &id)?, + None => U::default(), + }; + (update, result.context_table) + } else { + (U::default(), PluginContextTable::default()) + }; + let state = + hooks.post.then(|| CallState { runtime: Arc::clone(self), context_table, name: name.to_owned(), id }); + Ok((update, state)) + } + + async fn invoke( + &self, + hook: &'static str, + payload: MessagePayload, + context_table: Option, + ) -> PipelineResult { + let (result, background_tasks) = + self.manager.invoke_named::(hook, payload, Extensions::default(), context_table).await; + for error in &result.errors { + tracing::warn!( + hook, + plugin = error.plugin_name, + code = error.code.as_deref().unwrap_or(""), + proto_error_code = error.proto_error_code, + "CPEX plugin soft error" + ); + } + drop(background_tasks); + result + } +} + +impl CallState { + pub(crate) async fn after(&mut self, response: T) -> Result { + let result = self.invoke(&response).await?; + if result.is_denied() { + return Err(plugin_denied_error(T::OPERATION.subject(), result)); + } + self.apply(response, &result) + } + + pub(crate) async fn invoke(&mut self, response: &T) -> Result { + let payload = response.to_payload(&self.name, &self.id)?; + let result = self.runtime.invoke(T::OPERATION.hooks()[1], payload, Some(self.context_table.clone())).await; + if !result.is_denied() { + self.context_table = result.context_table.clone(); + } + Ok(result) + } + + pub(crate) fn apply(&self, response: T, result: &PipelineResult) -> Result { + match modified_message_payload(result) { + Some(payload) => response.apply_payload(payload, &self.name, &self.id), + None => Ok(response), + } + } } impl Drop for GatewayPluginRuntime { @@ -172,207 +181,3 @@ fn validate_gateway_supported_config(config: &CpexConfig) -> Result<(), GatewayP Ok(()) } - -impl GatewayPluginRuntime { - async fn invoke_cmf_hook( - &self, - hook_name: &'static str, - payload: MessagePayload, - context_table: Option, - ) -> PipelineResult { - let (result, background_tasks) = - self.manager.invoke_named::(hook_name, payload, Extensions::default(), context_table).await; - log_pipeline_errors(hook_name, &result); - drop(background_tasks); - result - } - - pub(crate) async fn before_tool_call( - &self, - request: &CallToolRequestParams, - tool_name: &str, - backend_name: &str, - ) -> Result { - if !self.hooks.tool.pre { - let state = self.hooks.tool.post.then(new_tool_call_state); - return Ok(ToolPreCallResult { arguments: ToolArgumentsUpdate::Unchanged, state }); - } - - let tool_call_id = next_tool_call_id(); - let original_payload = tool_call_payload(request, tool_name, backend_name, &tool_call_id); - let pre_result = self.invoke_cmf_hook(cmf_hook_names::TOOL_PRE_INVOKE, original_payload, None).await; - if pre_result.is_denied() { - return Err(plugin_denied_error("tool call", pre_result)); - } - - let arguments = effective_pre_args(request.arguments.as_ref(), &pre_result)?; - let state = Mutex::new(ToolCallState { context_table: pre_result.context_table, tool_call_id }); - Ok(ToolPreCallResult { arguments, state: Some(Arc::new(state)) }) - } - - 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_cmf_hook(cmf_hook_names::PROMPT_PRE_FETCH, payload, None).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 before_read_resource( - &self, - resource_uri: &str, - ) -> Result<(Option, Option), ErrorData> { - if !self.hooks.resource.pre && !self.hooks.resource.post { - return Ok((None, None)); - } - - let resource_request_id = next_resource_request_id(); - if !self.hooks.resource.pre { - return Ok(( - None, - Some(ResourceCallState { context_table: PluginContextTable::default(), resource_request_id }), - )); - } - - let payload = resource_request_payload(resource_uri, &resource_request_id); - let pre_result = self.invoke_cmf_hook(cmf_hook_names::RESOURCE_PRE_FETCH, payload, None).await; - if pre_result.is_denied() { - return Err(plugin_denied_error("resource", pre_result)); - } - let uri = effective_pre_resource_uri(&pre_result)?; - Ok(( - uri, - self.hooks - .resource - .post - .then_some(ResourceCallState { context_table: pre_result.context_table, resource_request_id }), - )) - } - - 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_cmf_hook(cmf_hook_names::PROMPT_POST_FETCH, 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_read_resource( - &self, - response: ReadResourceResult, - state: ResourceCallState, - ) -> Result { - let payload = resource_result_payload(&response, &state.resource_request_id) - .ok_or_else(|| ErrorData::internal_error("Resource response contains an unsupported content type", None))?; - let post_result = - self.invoke_cmf_hook(cmf_hook_names::RESOURCE_POST_FETCH, payload, Some(state.context_table)).await; - if post_result.is_denied() { - return Err(plugin_denied_error("resource", post_result)); - } - effective_post_resource_result(response, &post_result) - } - - pub(crate) async fn after_tool_call( - &self, - tool_name: &str, - response: CallToolResult, - state: Option, - ) -> Result { - if !self.hooks.tool.post { - return Ok(response); - } - - let state = state.and_then(|state| state.downcast::().ok()); - let Some(state) = state else { return Ok(response) }; - - let mut state = state.lock().await; - let post_result = self - .invoke_cmf_hook( - cmf_hook_names::TOOL_POST_INVOKE, - tool_result_payload(tool_name, &response, &state.tool_call_id), - Some(state.context_table.clone()), - ) - .await; - if post_result.is_denied() { - return Err(plugin_denied_error("tool call", post_result)); - } - - state.context_table = post_result.context_table.clone(); - Ok(effective_post_result(response, &post_result)) - } - - pub(crate) async fn after_tool_event( - &self, - tool_name: &str, - event: T, - state: Option, - ) -> Result, ErrorData> - where - T: Serialize + DeserializeOwned, - { - if !self.hooks.tool.post { - return Ok(Some(event)); - } - - let state = state.and_then(|state| state.downcast::().ok()); - let Some(state) = state else { return Ok(Some(event)) }; - - let content = serde_json::to_value(&event).unwrap_or(serde_json::Value::Null); - let mut state = state.lock().await; - let post_result = self - .invoke_cmf_hook( - cmf_hook_names::TOOL_POST_INVOKE, - tool_json_result_payload(tool_name, content, false, &state.tool_call_id), - Some(state.context_table.clone()), - ) - .await; - if post_result.is_denied() { - return Ok(None); - } - - state.context_table = post_result.context_table.clone(); - Ok(Some(effective_post_json(event, &post_result)?)) - } -} diff --git a/crates/contextforge-data-plane-cpex/src/tools/mod.rs b/crates/contextforge-data-plane-cpex/src/tools/mod.rs new file mode 100644 index 00000000..7731f875 --- /dev/null +++ b/crates/contextforge-data-plane-cpex/src/tools/mod.rs @@ -0,0 +1,187 @@ +use std::sync::Arc; + +use cpex::cpex_core::cmf::{ContentPart, MessagePayload, Role, ToolCall, ToolResult}; +use rmcp::{ + ErrorData, + model::{CallToolRequestParams, CallToolResult, ContentBlock}, + serde::{Serialize, de::DeserializeOwned}, +}; +use serde_json::{Map, Value}; +use tokio::sync::Mutex; + +use crate::{ + ArgumentsUpdate, GatewayPluginRuntimeHandle, PreHookResult, + cmf::{CmfResponse, Operation, message_payload}, + runtime::CallState, +}; + +fn tool_call_payload( + request: &CallToolRequestParams, + tool_name: &str, + backend_name: &str, + tool_call_id: &str, +) -> MessagePayload { + message_payload( + Role::Assistant, + vec![ContentPart::ToolCall { + content: ToolCall { + tool_call_id: tool_call_id.to_owned(), + name: tool_name.to_owned(), + arguments: request.arguments.clone().unwrap_or_default().into_iter().collect(), + namespace: Some(backend_name.to_owned()), + }, + }], + ) +} + +fn tool_result_payload(tool_name: &str, response: &CallToolResult, tool_call_id: &str) -> MessagePayload { + tool_json_result_payload( + tool_name, + serde_json::to_value(response).unwrap_or(Value::Null), + response.is_error.unwrap_or(false), + tool_call_id, + ) +} + +fn tool_json_result_payload(tool_name: &str, content: Value, is_error: bool, tool_call_id: &str) -> MessagePayload { + message_payload( + Role::Tool, + vec![ContentPart::ToolResult { + content: ToolResult { + tool_call_id: tool_call_id.to_owned(), + tool_name: tool_name.to_owned(), + content, + is_error, + }, + }], + ) +} + +fn tool_result_content(payload: &MessagePayload) -> Option { + payload.message.get_tool_results().first().map(|tool_result| tool_result.content.clone()) +} + +fn tool_call_arguments(payload: &MessagePayload) -> Option> { + payload + .message + .get_tool_calls() + .first() + .map(|tool_call| tool_call.arguments.clone().into_iter().collect::>()) +} + +fn tool_result_response(original: CallToolResult, payload: &MessagePayload) -> CallToolResult { + let mut result = payload.message.get_tool_results().first().map_or(original, |tool_result| { + serde_json::from_value::(tool_result.content.clone()).map_or_else( + |_| raw_tool_result(tool_result.content.clone(), tool_result.is_error), + |mut result| { + result.is_error = Some(tool_result.is_error); + result + }, + ) + }); + + let text = payload.message.get_text_content(); + if !text.is_empty() { + result.content.push(ContentBlock::text(text)); + } + + result +} + +fn raw_tool_result(value: Value, is_error: bool) -> CallToolResult { + match (value, is_error) { + (Value::String(text), false) => CallToolResult::success(vec![ContentBlock::text(text)]), + (Value::String(text), true) => CallToolResult::error(vec![ContentBlock::text(text)]), + (value, false) => CallToolResult::structured(value), + (value, true) => CallToolResult::structured_error(value), + } +} + +/// Shared by the final tool response and its progress notifications. +#[derive(Clone)] +pub struct ToolHookState(Arc>); + +impl std::fmt::Debug for ToolHookState { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("ToolHookState").finish_non_exhaustive() + } +} + +impl GatewayPluginRuntimeHandle { + pub async fn before_tool_call( + &self, + request: &CallToolRequestParams, + tool_name: &str, + backend_name: &str, + ) -> Result, ErrorData> { + let (arguments, state) = self + .current()? + .before( + Operation::Tool, + tool_name, + |id| tool_call_payload(request, tool_name, backend_name, id), + |payload, _| { + let arguments = tool_call_arguments(payload).ok_or_else(|| { + ErrorData::invalid_params("Plugin modified tool payload without a tool call", None) + })?; + Ok(ArgumentsUpdate::from_modified(request.arguments.as_ref(), arguments)) + }, + ) + .await?; + Ok(PreHookResult { arguments, state: state.map(|state| ToolHookState(Arc::new(Mutex::new(state)))) }) + } +} + +impl ToolHookState { + pub async fn after_tool_call(self, response: CallToolResult) -> Result { + self.0.lock().await.after(response).await + } + + /// Returns `None` when a plugin denies a progress or logging notification. + pub async fn after_stream_event(&self, event: T) -> Result, ErrorData> + where + T: Serialize + DeserializeOwned, + { + let mut state = self.0.lock().await; + let event = ToolEvent(event); + let result = state.invoke(&event).await?; + if result.is_denied() { + return Ok(None); + } + Ok(Some(state.apply(event, &result)?.0)) + } +} + +impl CmfResponse for CallToolResult { + const OPERATION: Operation = Operation::Tool; + + fn to_payload(&self, name: &str, id: &str) -> Result { + Ok(tool_result_payload(name, self, id)) + } + + fn apply_payload(self, payload: &MessagePayload, _name: &str, _id: &str) -> Result { + Ok(tool_result_response(self, payload)) + } +} + +struct ToolEvent(T); + +impl CmfResponse for ToolEvent { + const OPERATION: Operation = Operation::Tool; + + fn to_payload(&self, name: &str, id: &str) -> Result { + Ok(tool_json_result_payload(name, serde_json::to_value(&self.0).unwrap_or(Value::Null), false, id)) + } + + fn apply_payload(self, payload: &MessagePayload, _name: &str, _id: &str) -> Result { + let content = tool_result_content(payload).ok_or_else(|| { + ErrorData::invalid_params("Plugin modified stream event payload without a tool result", None) + })?; + serde_json::from_value(content).map(Self).map_err(|error| { + ErrorData::invalid_params(format!("Plugin modified stream event payload with invalid JSON: {error}"), None) + }) + } +} + +#[cfg(test)] +mod tests; diff --git a/crates/contextforge-data-plane-cpex/src/tools/tests.rs b/crates/contextforge-data-plane-cpex/src/tools/tests.rs new file mode 100644 index 00000000..ea7ce231 --- /dev/null +++ b/crates/contextforge-data-plane-cpex/src/tools/tests.rs @@ -0,0 +1,16 @@ +use super::*; + +#[test] +fn tool_result_response_uses_cmf_error_flag_for_nested_mcp_result() { + let original = CallToolResult::success(vec![ContentBlock::text("original")]); + let nested = CallToolResult::success(vec![ContentBlock::text("changed")]); + let mut payload = tool_result_payload("sum", &nested, "call-1"); + let ContentPart::ToolResult { content } = &mut payload.message.content[0] else { + panic!("expected tool result"); + }; + content.is_error = true; + + let result = tool_result_response(original, &payload); + + assert_eq!(Some(true), result.is_error); +} diff --git a/crates/contextforge-data-plane-lib/src/gateway/backend_client.rs b/crates/contextforge-data-plane-lib/src/gateway/backend_client.rs index 8c2cf3f9..e964d180 100644 --- a/crates/contextforge-data-plane-lib/src/gateway/backend_client.rs +++ b/crates/contextforge-data-plane-lib/src/gateway/backend_client.rs @@ -1,6 +1,6 @@ use std::{collections::HashMap, sync::Arc}; -use contextforge_data_plane_cpex::{GatewayPluginRuntimeHandle, RuntimeHookState}; +use contextforge_data_plane_cpex::ToolHookState; use rmcp::{ ClientHandler, Peer, RoleClient, RoleServer, model::{ @@ -17,24 +17,19 @@ use tracing::{debug, warn}; #[derive(Clone)] pub(crate) struct GatewayBackendClient { initialize_request: InitializeRequestParams, - plugin_runtime: Option, in_flight_calls: Arc>>>, } #[derive(Debug)] struct InFlightToolCall { downstream_progress_token: ProgressToken, - tool_name: String, - post_state: Option, + post_state: Option, downstream: Peer, } impl GatewayBackendClient { - pub(crate) fn new( - initialize_request: InitializeRequestParams, - plugin_runtime: Option, - ) -> Self { - Self { initialize_request, plugin_runtime, in_flight_calls: Arc::default() } + pub(crate) fn new(initialize_request: InitializeRequestParams) -> Self { + Self { initialize_request, in_flight_calls: Arc::default() } } /// Starts a backend tool call while preventing an immediate progress @@ -45,11 +40,10 @@ impl GatewayBackendClient { peer: &Peer, request: CallToolRequestParams, downstream_progress_token: Option, - tool_name: String, downstream: Peer, - post_state: Option, + post_state: Option, ) -> Result, ServiceError> { - debug!("track_tool_call {tool_name} {downstream_progress_token:?} {post_state:?}"); + debug!("track_tool_call {downstream_progress_token:?} {post_state:?}"); let request = ClientRequest::CallToolRequest(Request::new(request)); let Some(downstream_progress_token) = downstream_progress_token else { return peer.send_cancellable_request(request, PeerRequestOptions::no_options()).await; @@ -61,7 +55,7 @@ impl GatewayBackendClient { let mut calls = self.in_flight_calls.write().await; let handle = peer.send_cancellable_request(request, PeerRequestOptions::no_options()).await?; let backend_progress_token = handle.progress_token.clone(); - let call = Arc::new(InFlightToolCall { downstream_progress_token, tool_name, post_state, downstream }); + let call = Arc::new(InFlightToolCall { downstream_progress_token, post_state, downstream }); calls.insert(backend_progress_token, call); Ok(handle) } @@ -81,10 +75,10 @@ impl GatewayBackendClient { where T: Serialize + DeserializeOwned, { - let Some(plugin_runtime) = &self.plugin_runtime else { + let Some(state) = &call.post_state else { return Some(event); }; - match plugin_runtime.after_stream_event(&call.tool_name, event, call.post_state.clone()).await { + match state.after_stream_event(event).await { Ok(event) => event, Err(error) => { warn!("call_tool: plugin rejected backend notification: {error:?}"); diff --git a/crates/contextforge-data-plane-lib/src/gateway/mcp_service/initialization.rs b/crates/contextforge-data-plane-lib/src/gateway/mcp_service/initialization.rs index 9096475c..f2d3fec2 100644 --- a/crates/contextforge-data-plane-lib/src/gateway/mcp_service/initialization.rs +++ b/crates/contextforge-data-plane-lib/src/gateway/mcp_service/initialization.rs @@ -64,7 +64,7 @@ pub(super) async fn connect_backend_for_request( ) .with_protocol_version(backend.mcp_protocol_version.clone()); - let backend_client = GatewayBackendClient::new(client_info, mcp_service.plugin_runtime.clone()); + let backend_client = GatewayBackendClient::new(client_info); serve_client_with_lifecycle_and_ct( backend_client, 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 f092bfb7..0260ce7f 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,4 +1,4 @@ -use contextforge_data_plane_cpex::PromptPreFetchResult; +use contextforge_data_plane_cpex::PreHookResult; use rmcp::{ ErrorData, RoleServer, model::{ErrorCode, GetPromptRequestParams, GetPromptResponse}, @@ -38,21 +38,21 @@ pub(super) async fn get_prompt( let pre_result = if let Some(plugin_runtime) = &mcp_service.plugin_runtime { plugin_runtime.before_get_prompt(&request, &prompt_name, &backend_name).await? } else { - PromptPreFetchResult::unchanged() + PreHookResult::default() }; let mut backend_service = connect_backend_for_request(mcp_service, &backend_name, backend, &cx).await?; let mut routed_request = request; - pre_result.arguments.apply_to_request(&mut routed_request, &prompt_name); + routed_request.name.clone_from(&prompt_name); + pre_result.arguments.apply_to(&mut routed_request.arguments); let response = backend_service.get_prompt(routed_request).await; if let Err(error) = backend_service.close().await { tracing::warn!("get_prompt: backend cleanup failed backend_name = {backend_name} error = {error:?}"); } let response = response.map_err(|error| backend_forward_error("get_prompt", &backend_name, &error))?; info!("get_prompt: backend {backend_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 + let response = match pre_result.state { + Some(state) => state.after_get_prompt(response).await?, + None => response, }; Ok(response.into()) } diff --git a/crates/contextforge-data-plane-lib/src/gateway/mcp_service/tools.rs b/crates/contextforge-data-plane-lib/src/gateway/mcp_service/tools.rs index abbe8aad..73ce227b 100644 --- a/crates/contextforge-data-plane-lib/src/gateway/mcp_service/tools.rs +++ b/crates/contextforge-data-plane-lib/src/gateway/mcp_service/tools.rs @@ -1,4 +1,4 @@ -use contextforge_data_plane_cpex::ToolPreCallResult; +use contextforge_data_plane_cpex::PreHookResult; use http::request::Parts; use rmcp::{ ErrorData, RoleServer, @@ -55,24 +55,18 @@ pub(super) async fn call_tool( let pre_result = if let Some(plugin_runtime) = &mcp_service.plugin_runtime { plugin_runtime.before_tool_call(&request, &tool_name, &backend_name).await? } else { - ToolPreCallResult::unchanged() + PreHookResult::default() }; let mut backend_service = connect_backend_for_request(mcp_service, &backend_name, backend, &cx).await?; let post_state = pre_result.state; let mut routed_request = request; - pre_result.arguments.apply_to_request(&mut routed_request, &tool_name); + routed_request.name = tool_name.clone().into(); + pre_result.arguments.apply_to(&mut routed_request.arguments); let progress_token = cx.meta.get_progress_token(); let handle = backend_service .service() - .start_tool_call( - backend_service.peer(), - routed_request, - progress_token, - tool_name.clone(), - cx.peer.clone(), - post_state.clone(), - ) + .start_tool_call(backend_service.peer(), routed_request, progress_token, cx.peer.clone(), post_state.clone()) .await .map_err(|error| backend_forward_error("call_tool", &backend_name, &error))?; let backend_progress_token = handle.progress_token.clone(); @@ -83,11 +77,9 @@ pub(super) async fn call_tool( } let response = response.map_err(|error| backend_forward_error("call_tool", &backend_name, &error))?; - let response = match (&mcp_service.plugin_runtime, post_state) { - (Some(plugin_runtime), Some(post_state)) => { - plugin_runtime.after_tool_call(&tool_name, response, Some(post_state)).await? - }, - _ => response, + let response = match post_state { + Some(state) => state.after_tool_call(response).await?, + None => response, }; info!("call_tool: backend {backend_name} completed"); Ok(response.into())