From e4193e92c6a0ebd05a02bfd1991b11d6d70e32cf Mon Sep 17 00:00:00 2001 From: lucarlig Date: Thu, 3 Sep 2026 10:50:40 +0100 Subject: [PATCH 1/3] feat: run CPEX hooks for resource reads Signed-off-by: lucarlig --- .secrets.baseline | 4 +- README.md | 4 +- _context/wiki/architecture.md | 6 +- _context/wiki/config.md | 8 +- .../contextforge-data-plane-cpex/src/cmf.rs | 157 +++++++++++++++++- .../src/factory.rs | 2 + .../src/handle.rs | 46 ++++- .../contextforge-data-plane-cpex/src/hooks.rs | 10 ++ .../contextforge-data-plane-cpex/src/lib.rs | 4 +- .../src/pipeline.rs | 33 +++- .../src/runtime.rs | 103 +++++++++++- .../src/gateway/mcp_service/resources.rs | 11 ++ .../tests/gateway_plugins.rs | 40 ++++- .../tests/support/mod.rs | 4 +- .../tests/support/plugin_gateway.rs | 27 ++- .../plugins/cpex-secrets-detection/README.md | 10 +- .../plugins/cpex-secrets-detection/src/lib.rs | 3 +- 17 files changed, 427 insertions(+), 45 deletions(-) diff --git a/.secrets.baseline b/.secrets.baseline index 6eeb6e16..04d8e0b0 100644 --- a/.secrets.baseline +++ b/.secrets.baseline @@ -3,7 +3,7 @@ "files": "(?x)(Cargo\\.lock$|\\.lock$)|^\\.secrets\\.baseline$|^.secrets.baseline$", "lines": null }, - "generated_at": "2026-08-26T14:21:23Z", + "generated_at": "2026-09-03T09:50:29Z", "plugins_used": [ { "name": "AWSKeyDetector" @@ -208,7 +208,7 @@ "hashed_secret": "86de8c52637ec530fe39b0a8471da9b8764d5242", "is_secret": false, "is_verified": false, - "line_number": 610, + "line_number": 609, "type": "AWS Access Key", "verified_result": null } diff --git a/README.md b/README.md index 814356db..8dad5f3d 100644 --- a/README.md +++ b/README.md @@ -60,8 +60,8 @@ Activation requires all three pieces: - Runtime flag: `--runtime-plugins-enabled true` - Redis config key: `ContextForgeGatewayRuntimePluginConfig` -The plugin kind is `validator/secrets-detection`. The data plane currently -wires only `cmf.tool_pre_invoke` and `cmf.tool_post_invoke`. +The plugin kind is `validator/secrets-detection`. The dataplane wires CMF hooks +for tool calls, prompt fetches, and resource reads. Example run command: diff --git a/_context/wiki/architecture.md b/_context/wiki/architecture.md index 593e7c72..449ce01b 100644 --- a/_context/wiki/architecture.md +++ b/_context/wiki/architecture.md @@ -154,7 +154,7 @@ The binary sets `tikv_jemallocator` as the global allocator. jemalloc holds up b - `initialize` opens one backend transport per configured backend concurrently (`futures::future::join_all`); a failed backend degrades that backend only. - List methods fan out to all connected backends concurrently and merge. - Targeted calls (except `call_tool`) resolve exactly one backend service handle from `BackendTransports`. -- `call_tool` creates a fresh per-request backend connection via `connect_backend_for_request`, runs pre/post plugin hooks, executes the call, then explicitly closes the connection before returning. +- Targeted tool, prompt, and resource calls run configured pre/post plugin hooks after backend routing. `call_tool` creates a fresh per-request backend connection via `connect_backend_for_request`, then explicitly closes it before returning. - `call_tool` watches the downstream cancellation token and forwards a cancel to the backend if the client gives up first; backend progress notifications are forwarded downstream while the call is in flight. ## Startup And Response Flow @@ -180,7 +180,7 @@ Response unwind order (Tower layers execute outside-in, so unwind is inside-out) ```text backend response - -> response plugin hooks (call_tool only) + -> response plugin hooks (tool, prompt, and resource calls) -> merge / namespace / pass through -> virtual_host_config_layer response side -> user_config_store_layer response side @@ -227,7 +227,7 @@ Do not bury transport security decisions inside MCP method handlers. They belong ## Plugin Hook Expansion Requirements -Current supported hooks are intentionally narrow (`cmf.tool_pre_invoke`, `cmf.tool_post_invoke`). Before adding any new hook point, define all of the following: +Current supported hooks cover tool, prompt, and resource pre/post lifecycles. Before adding any new hook point, define all of the following: | Requirement | Why | | --- | --- | diff --git a/_context/wiki/config.md b/_context/wiki/config.md index 315ef53d..6059df2b 100644 --- a/_context/wiki/config.md +++ b/_context/wiki/config.md @@ -174,8 +174,8 @@ RuntimePluginConfigDocument cpex: CpexConfig ``` -Supported: `cmf.tool_pre_invoke`, `cmf.tool_post_invoke`, `cmf.prompt_pre_fetch`, `cmf.prompt_post_fetch` only. -Rejected: routing-based selection, plugin dirs, global policies, resource and LLM hooks, plugin conditions. +Supported: tool, prompt, and resource pre/post CMF hooks. +Rejected: routing-based selection, plugin dirs, global policies, LLM hooks, plugin conditions. Config validation and `CmfPluginFactory` registration must agree on that list: a hook accepted by validation but not registered leaves the plugin loaded and silently inert. Reload watcher: 10-minute interval. Invalid reload → runtime marked failed. @@ -203,6 +203,10 @@ MCP prompt results carry no error flag, so a plugin setting `is_error` on the CM Binary resource blobs reach plugins by URI and MIME type but not by content: CMF stores decoded bytes while MCP sends base64. A plugin can deny such a message; editing one fails the write-back. +### Resource Read Hook Behavior + +For `resources/read`, the pre hook runs after routing and may allow or deny the canonical backend-local URI, but cannot change it. The post hook can redact text resources or deny the response. URI, MIME type, item count, and blob identity must remain stable; unsupported or lossy edits fail closed. + ### Demo Plugin Workflow The optional `test-plugins` feature compiles demo factories from the `cpex-plugins-rs` repository. Redis configuration activates factories already present in the binary; it never loads new Rust code into a running process. diff --git a/crates/contextforge-data-plane-cpex/src/cmf.rs b/crates/contextforge-data-plane-cpex/src/cmf.rs index eeb74c97..5f1a66a8 100644 --- a/crates/contextforge-data-plane-cpex/src/cmf.rs +++ b/crates/contextforge-data-plane-cpex/src/cmf.rs @@ -6,7 +6,7 @@ use cpex::cpex_core::cmf::{ }; use rmcp::model::{ CallToolRequestParams, CallToolResult, ContentBlock, GetPromptRequestParams, GetPromptResult, PromptMessage, - Resource as McpResource, ResourceContents, Role as McpRole, + ReadResourceResult, Resource as McpResource, ResourceContents, Role as McpRole, }; use serde_json::{Map, Value}; @@ -33,6 +33,119 @@ pub(crate) fn tool_call_payload( } } +pub(crate) fn resource_request_payload(resource_uri: &str, resource_request_id: &str) -> MessagePayload { + MessagePayload { + message: Message { + schema_version: "2.0".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_request_matches( + payload: &MessagePayload, + resource_uri: &str, + resource_request_id: &str, +) -> bool { + let [ContentPart::ResourceRef { content }] = payload.message.content.as_slice() else { + return false; + }; + payload.message.role == Role::User + && content.resource_request_id == resource_request_id + && content.uri == resource_uri + && matches!(content.resource_type, ResourceType::Uri) + && content.name.is_none() + && content.range_start.is_none() + && content.range_end.is_none() + && content.selector.is_none() +} + +pub(crate) fn resource_result_payload( + response: &ReadResourceResult, + resource_request_id: &str, +) -> Option { + let content = response + .contents + .iter() + .map(|content| { + let (uri, mime_type, text) = match content { + 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, + }; + Some(ContentPart::Resource { + content: CmfResource { + resource_request_id: resource_request_id.to_owned(), + uri, + resource_type: ResourceType::Uri, + content: text, + mime_type, + ..Default::default() + }, + }) + }) + .collect::>>()?; + Some(MessagePayload { + message: Message { schema_version: "2.0".to_owned(), role: Role::Assistant, content, channel: None }, + }) +} + +pub(crate) fn resource_result_response( + mut original: ReadResourceResult, + payload: &MessagePayload, + resource_request_id: &str, +) -> Option { + if payload.message.role != Role::Assistant || 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; + }; + if content.resource_request_id != resource_request_id + || !matches!(content.resource_type, ResourceType::Uri) + || content.name.is_some() + || content.description.is_some() + || content.blob.is_some() + || content.size_bytes.is_some() + || !content.annotations.is_empty() + || content.version.is_some() + { + return None; + } + match original { + ResourceContents::TextResourceContents { uri, mime_type, text, .. } => { + if content.uri != *uri || content.mime_type != *mime_type { + return None; + } + *text = content.content.clone()?; + }, + ResourceContents::BlobResourceContents { uri, mime_type, .. } => { + if content.uri != *uri || content.mime_type != *mime_type || content.content.is_some() { + return None; + } + }, + _ => 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, @@ -347,6 +460,48 @@ fn mcp_prompt_message(message: &Message) -> Option { mod tests { use super::*; + #[test] + fn resource_request_rejects_uri_mutation() { + let mut payload = resource_request_payload("file:///password.env", "resource-1"); + let ContentPart::ResourceRef { content } = &mut payload.message.content[0] else { + panic!("expected resource reference"); + }; + content.uri = "file:///other.env".to_owned(); + + assert!(!resource_request_matches(&payload, "file:///password.env", "resource-1")); + } + + #[test] + fn resource_result_response_applies_only_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, "resource-1").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_rejects_uri_mutation() { + let original = ReadResourceResult::new(vec![ResourceContents::text("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.uri = "file:///other.env".to_owned(); + + assert!(resource_result_response(original, &payload, "resource-1").is_none()); + } + fn text_prompt() -> GetPromptResult { GetPromptResult::new(vec![PromptMessage::new_text(McpRole::User, "review of weather")]) } diff --git a/crates/contextforge-data-plane-cpex/src/factory.rs b/crates/contextforge-data-plane-cpex/src/factory.rs index 598d8a2e..a2c8c6e3 100644 --- a/crates/contextforge-data-plane-cpex/src/factory.rs +++ b/crates/contextforge-data-plane-cpex/src/factory.rs @@ -51,6 +51,8 @@ fn cmf_hook_name(hook: &str) -> Option<&'static str> { 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, } } diff --git a/crates/contextforge-data-plane-cpex/src/handle.rs b/crates/contextforge-data-plane-cpex/src/handle.rs index a2f504a7..6dcd0dd8 100644 --- a/crates/contextforge-data-plane-cpex/src/handle.rs +++ b/crates/contextforge-data-plane-cpex/src/handle.rs @@ -13,7 +13,9 @@ use cpex::cpex_core::{ }; use rmcp::{ ErrorData, - model::{CallToolRequestParams, CallToolResult, ErrorCode, GetPromptRequestParams, GetPromptResult}, + model::{ + CallToolRequestParams, CallToolResult, ErrorCode, GetPromptRequestParams, GetPromptResult, ReadResourceResult, + }, serde::{Serialize, de::DeserializeOwned}, }; use tokio::task::JoinHandle; @@ -21,7 +23,7 @@ use tokio::task::JoinHandle; use crate::{ config::{LoadedRuntimePluginConfig, RedisRuntimePluginConfigStore, RuntimePluginConfigStore, cpex_config}, error::GatewayPluginRuntimeError, - hooks::{PromptPreFetchResult, RuntimeHookError, RuntimeHookState, ToolPreCallResult}, + hooks::{PromptPreFetchResult, ResourcePreFetchResult, RuntimeHookError, RuntimeHookState, ToolPreCallResult}, runtime::GatewayPluginRuntime, }; @@ -280,6 +282,21 @@ impl GatewayPluginRuntimeHandle { 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 mut result = runtime.before_read_resource(resource_uri).await?; + if runtime.has_resource_post_hook() { + let state = result.state.take(); + result.state = Some(Arc::new(RegistryCallState { runtime: Arc::clone(runtime), state })); + } else { + result.state = None; + } + Ok(result) + } + pub async fn after_get_prompt( &self, prompt_name: &str, @@ -292,6 +309,17 @@ impl GatewayPluginRuntimeHandle { } } + pub async fn after_read_resource( + &self, + response: ReadResourceResult, + state: Option, + ) -> Result { + match state.and_then(|state| state.downcast::().ok()) { + Some(state) => state.runtime.after_read_resource(response, state.state.clone()).await, + None => Ok(response), + } + } + pub async fn after_tool_call( &self, tool_name: &str, @@ -324,7 +352,7 @@ impl GatewayPluginRuntimeHandle { fn runtime_failed_error(state: &RuntimeState) -> ErrorData { if let RuntimeState::Failed(error) = state { - tracing::warn!(%error, "rejecting tool call because CPEX runtime is failed"); + tracing::warn!(%error, "rejecting MCP call because CPEX runtime is failed"); } ErrorData { code: ErrorCode::INTERNAL_ERROR, message: "Runtime plugin reload failed".into(), data: None } } @@ -640,6 +668,7 @@ mod tests { 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, _ => return None, }; Some(( @@ -769,6 +798,17 @@ mod tests { 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 runtime_config_loads_registered_factory_plugin() { let plugin = diff --git a/crates/contextforge-data-plane-cpex/src/hooks.rs b/crates/contextforge-data-plane-cpex/src/hooks.rs index b6a9a63a..0b481601 100644 --- a/crates/contextforge-data-plane-cpex/src/hooks.rs +++ b/crates/contextforge-data-plane-cpex/src/hooks.rs @@ -52,6 +52,16 @@ pub struct PromptPreFetchResult { pub state: Option, } +pub struct ResourcePreFetchResult { + pub state: Option, +} + +impl ResourcePreFetchResult { + pub fn unchanged() -> Self { + Self { state: None } + } +} + impl PromptPreFetchResult { pub fn unchanged() -> Self { Self { arguments: PromptArgumentsUpdate::Unchanged, state: None } diff --git a/crates/contextforge-data-plane-cpex/src/lib.rs b/crates/contextforge-data-plane-cpex/src/lib.rs index 5d1be674..944bfb22 100644 --- a/crates/contextforge-data-plane-cpex/src/lib.rs +++ b/crates/contextforge-data-plane-cpex/src/lib.rs @@ -11,6 +11,6 @@ pub use error::GatewayPluginRuntimeError; pub use factory::CmfPluginFactory; pub use handle::{CpexRuntimeRegistry, GatewayPluginRuntimeHandle}; pub use hooks::{ - PromptArgumentsUpdate, PromptPreFetchResult, RuntimeHookError, RuntimeHookState, ToolArgumentsUpdate, - ToolPreCallResult, + PromptArgumentsUpdate, PromptPreFetchResult, ResourcePreFetchResult, RuntimeHookError, RuntimeHookState, + ToolArgumentsUpdate, ToolPreCallResult, }; diff --git a/crates/contextforge-data-plane-cpex/src/pipeline.rs b/crates/contextforge-data-plane-cpex/src/pipeline.rs index 9e4ecab0..a03081ca 100644 --- a/crates/contextforge-data-plane-cpex/src/pipeline.rs +++ b/crates/contextforge-data-plane-cpex/src/pipeline.rs @@ -2,7 +2,7 @@ use cpex::cpex_core::cmf::MessagePayload; use cpex::cpex_core::executor::PipelineResult; use rmcp::{ ErrorData, - model::{CallToolResult, ErrorCode, GetPromptResult}, + model::{CallToolResult, ErrorCode, GetPromptResult, ReadResourceResult}, serde::de::DeserializeOwned, }; use tracing::warn; @@ -10,8 +10,8 @@ use tracing::warn; use crate::{ PromptArgumentsUpdate, ToolArgumentsUpdate, cmf::{ - prompt_request_arguments, prompt_result_rejection, prompt_result_response, tool_call_arguments, - tool_result_content, tool_result_response, + prompt_request_arguments, prompt_result_rejection, prompt_result_response, resource_request_matches, + resource_result_response, tool_call_arguments, tool_result_content, tool_result_response, }, }; @@ -97,6 +97,33 @@ pub(crate) fn effective_post_prompt_result( }) } +pub(crate) fn validate_pre_resource_result( + result: &PipelineResult, + resource_uri: &str, + resource_request_id: &str, +) -> Result<(), ErrorData> { + let Some(payload) = modified_message_payload(result) else { + return Ok(()); + }; + if resource_request_matches(payload, resource_uri, resource_request_id) { + Ok(()) + } else { + Err(ErrorData::internal_error("Plugin attempted to modify the canonical resource route", None)) + } +} + +pub(crate) fn effective_post_resource_result( + original: ReadResourceResult, + result: &PipelineResult, + resource_request_id: &str, +) -> Result { + let Some(payload) = modified_message_payload(result) else { + return Ok(original); + }; + resource_result_response(original, payload, resource_request_id) + .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, diff --git a/crates/contextforge-data-plane-cpex/src/runtime.rs b/crates/contextforge-data-plane-cpex/src/runtime.rs index 36c2efb7..fd9d0ce6 100644 --- a/crates/contextforge-data-plane-cpex/src/runtime.rs +++ b/crates/contextforge-data-plane-cpex/src/runtime.rs @@ -14,20 +14,22 @@ use cpex::cpex_core::{ }; use rmcp::{ ErrorData, - model::{CallToolRequestParams, CallToolResult, GetPromptRequestParams, GetPromptResult}, + model::{CallToolRequestParams, CallToolResult, GetPromptRequestParams, GetPromptResult, ReadResourceResult}, serde::{Serialize, de::DeserializeOwned}, }; use tokio::sync::Mutex; use crate::{ cmf::{ - prompt_request_payload, prompt_result_payload, tool_call_payload, tool_json_result_payload, tool_result_payload, + prompt_request_payload, prompt_result_payload, resource_request_payload, resource_result_payload, + tool_call_payload, tool_json_result_payload, tool_result_payload, }, error::GatewayPluginRuntimeError, - hooks::{PromptPreFetchResult, RuntimeHookState, ToolArgumentsUpdate, ToolPreCallResult}, + hooks::{PromptPreFetchResult, ResourcePreFetchResult, RuntimeHookState, ToolArgumentsUpdate, ToolPreCallResult}, pipeline::{ - effective_post_json, effective_post_prompt_result, effective_post_result, effective_pre_args, - effective_pre_prompt_args, log_pipeline_errors, plugin_denied_error, + effective_post_json, effective_post_prompt_result, effective_post_resource_result, effective_post_result, + effective_pre_args, effective_pre_prompt_args, log_pipeline_errors, plugin_denied_error, + validate_pre_resource_result, }, }; @@ -41,6 +43,7 @@ struct HookPair { struct HookPresence { tool: HookPair, prompt: HookPair, + resource: HookPair, } #[derive(Default)] @@ -82,6 +85,19 @@ fn new_prompt_call_state(context_table: PluginContextTable, prompt_request_id: S Arc::new(PromptCallState { context_table, prompt_request_id }) } +struct ResourceCallState { + context_table: Option, + resource_request_id: String, +} + +fn next_resource_request_id() -> String { + format!("gateway-resource-request-{}", CORRELATION_ID.fetch_add(1, Ordering::Relaxed)) +} + +fn new_resource_call_state(context_table: Option, resource_request_id: String) -> RuntimeHookState { + Arc::new(ResourceCallState { context_table, resource_request_id }) +} + impl GatewayPluginRuntime { pub(crate) fn has_post_hook(&self) -> bool { self.hooks.tool.post @@ -91,6 +107,10 @@ impl GatewayPluginRuntime { self.hooks.prompt.post } + pub(crate) fn has_resource_post_hook(&self) -> bool { + self.hooks.resource.post + } + pub(crate) async fn from_config( config: CpexConfig, factories: &PluginFactoryRegistry, @@ -106,6 +126,10 @@ impl GatewayPluginRuntime { 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 manager = PluginManager::from_config(config, factories) .map_err(|source| GatewayPluginRuntimeError::Configuration { hook: "config", source })?; @@ -128,11 +152,13 @@ impl Drop for GatewayPluginRuntime { } } -const SUPPORTED_HOOKS: [&str; 4] = [ +const SUPPORTED_HOOKS: [&str; 6] = [ cmf_hook_names::TOOL_PRE_INVOKE, cmf_hook_names::TOOL_POST_INVOKE, cmf_hook_names::PROMPT_PRE_FETCH, cmf_hook_names::PROMPT_POST_FETCH, + cmf_hook_names::RESOURCE_PRE_FETCH, + cmf_hook_names::RESOURCE_POST_FETCH, ]; fn declares(config: &CpexConfig, hook_name: &str) -> bool { @@ -270,6 +296,51 @@ impl GatewayPluginRuntime { Ok(PromptPreFetchResult { arguments, state }) } + async fn invoke_resource_pre(&self, payload: MessagePayload) -> PipelineResult { + let (result, background_tasks) = self + .manager + .invoke_named::(cmf_hook_names::RESOURCE_PRE_FETCH, payload, Extensions::default(), None) + .await; + log_pipeline_errors(cmf_hook_names::RESOURCE_PRE_FETCH, &result); + drop(background_tasks); + result + } + + async fn invoke_resource_post( + &self, + payload: MessagePayload, + context_table: Option, + ) -> PipelineResult { + let (result, background_tasks) = self + .manager + .invoke_named::(cmf_hook_names::RESOURCE_POST_FETCH, payload, Extensions::default(), context_table) + .await; + log_pipeline_errors(cmf_hook_names::RESOURCE_POST_FETCH, &result); + drop(background_tasks); + result + } + + pub(crate) async fn before_read_resource(&self, resource_uri: &str) -> Result { + let resource_request_id = next_resource_request_id(); + if !self.hooks.resource.pre { + let state = self.hooks.resource.post.then(|| new_resource_call_state(None, resource_request_id)); + return Ok(ResourcePreFetchResult { state }); + } + + let payload = resource_request_payload(resource_uri, &resource_request_id); + let pre_result = self.invoke_resource_pre(payload).await; + if pre_result.is_denied() { + return Err(plugin_denied_error("resource", pre_result)); + } + validate_pre_resource_result(&pre_result, resource_uri, &resource_request_id)?; + let state = self + .hooks + .resource + .post + .then(|| new_resource_call_state(Some(pre_result.context_table), resource_request_id)); + Ok(ResourcePreFetchResult { state }) + } + pub(crate) async fn after_get_prompt( &self, prompt_name: &str, @@ -292,6 +363,26 @@ impl GatewayPluginRuntime { effective_post_prompt_result(response, &post_result, prompt_name, &state.prompt_request_id) } + pub(crate) async fn after_read_resource( + &self, + response: ReadResourceResult, + state: Option, + ) -> Result { + if !self.hooks.resource.post { + return Ok(response); + } + + let state = state.and_then(|state| state.downcast::().ok()); + let Some(state) = state else { return Ok(response) }; + 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_resource_post(payload, state.context_table.clone()).await; + if post_result.is_denied() { + return Err(plugin_denied_error("resource", post_result)); + } + effective_post_resource_result(response, &post_result, &state.resource_request_id) + } + pub(crate) async fn after_tool_call( &self, tool_name: &str, diff --git a/crates/contextforge-data-plane-lib/src/gateway/mcp_service/resources.rs b/crates/contextforge-data-plane-lib/src/gateway/mcp_service/resources.rs index b4e02799..a2b5bab1 100644 --- a/crates/contextforge-data-plane-lib/src/gateway/mcp_service/resources.rs +++ b/crates/contextforge-data-plane-lib/src/gateway/mcp_service/resources.rs @@ -1,3 +1,4 @@ +use contextforge_data_plane_cpex::ResourcePreFetchResult; use rmcp::{ ErrorData, RoleServer, model::{ @@ -45,6 +46,11 @@ pub(super) async fn read_resource( })?; let service_name = backend_name.clone(); + let pre_result = if let Some(plugin_runtime) = &mcp_service.plugin_runtime { + plugin_runtime.before_read_resource(&resource_uri).await? + } else { + ResourcePreFetchResult::unchanged() + }; let mut backend_service = connect_backend_for_request(mcp_service, &backend_name, backend, &cx).await?; let mut routed_request = request; @@ -55,6 +61,11 @@ pub(super) async fn read_resource( tracing::warn!("read_resource: backend cleanup failed backend_name = {service_name} error = {error:?}"); } let response = response.map_err(|error| backend_forward_error("read_resource", &service_name, &error))?; + let response = if let Some(plugin_runtime) = &mcp_service.plugin_runtime { + plugin_runtime.after_read_resource(response, pre_result.state).await? + } else { + response + }; info!("read_resource: backend {service_name} returned {} contents", response.contents.len()); diff --git a/crates/contextforge-data-plane-lib/tests/gateway_plugins.rs b/crates/contextforge-data-plane-lib/tests/gateway_plugins.rs index 7be3d352..0a1f78e3 100644 --- a/crates/contextforge-data-plane-lib/tests/gateway_plugins.rs +++ b/crates/contextforge-data-plane-lib/tests/gateway_plugins.rs @@ -11,18 +11,18 @@ use rmcp::{ model::{ CallToolRequestParams, CallToolResult, ClientCapabilities, ClientRequest, ContentBlock, ErrorCode, GetPromptRequestParams, GetPromptResult, Implementation, InitializeRequestParams, ProgressNotificationParam, - Request, ResourceContents, Role as McpRole, ServerResult, + ReadResourceRequestParams, Request, ResourceContents, Role as McpRole, ServerResult, }, service::{NotificationContext, PeerRequestOptions, RequestHandle, RoleClient, RunningService}, }; use serde_json::{Map, Value, json}; use support::{ - BACKEND_PROMPT_IMAGE, BACKEND_PROMPT_RESOURCE, POST_DENY_ERROR_CODE, PRE_DENY_ERROR_CODE, PROMPT_ERROR_MESSAGE, - PROMPT_POST_DENY_ERROR_CODE, PromptBehavior, PromptTestPlugin, REWRITTEN_PROMPT_RESOURCE, REWRITTEN_PROMPT_TEXT, - REWRITTEN_PROMPT_TOPIC, REWRITTEN_SUM_A, REWRITTEN_SUM_B, RunningGateway, TEST_USER_ID, TestPlugin, error_code, - error_parts, runtime_with_post, runtime_with_pre, runtime_with_pre_and_post, runtime_with_prompt_plugin, - start_gateway, start_gateway_with_events, start_gateway_with_json_backend_responses, + BACKEND_PROMPT_IMAGE, BACKEND_PROMPT_RESOURCE, BACKEND_RESOURCE_SECRET, POST_DENY_ERROR_CODE, PRE_DENY_ERROR_CODE, + PROMPT_ERROR_MESSAGE, PROMPT_POST_DENY_ERROR_CODE, PromptBehavior, PromptTestPlugin, REWRITTEN_PROMPT_RESOURCE, + REWRITTEN_PROMPT_TEXT, REWRITTEN_PROMPT_TOPIC, REWRITTEN_SUM_A, REWRITTEN_SUM_B, RunningGateway, TEST_USER_ID, + TestPlugin, error_code, error_parts, runtime_with_post, runtime_with_pre, runtime_with_pre_and_post, + runtime_with_prompt_plugin, start_gateway, start_gateway_with_events, start_gateway_with_json_backend_responses, start_gateway_with_parameter_headers, sum_request, text, token, }; @@ -667,6 +667,34 @@ async fn secrets_detection_pre_hook_respects_field_allowlist() { assert_eq!(Some(&Value::from(ignored_secret)), args.get("ignored")); } +#[tokio::test(flavor = "multi_thread", worker_threads = 1)] +async fn secrets_detection_resource_post_hook_redacts_password_resource() { + let runtime = runtime_with_secrets_detection( + vec![cmf_hook_names::RESOURCE_POST_FETCH], + json!({ + "redact": true, + "redaction_text": "[redacted]", + "block_on_detection": false, + }), + ) + .await; + let gateway = start_gateway(TEST_USER_ID, true, runtime).await; + let service = gateway.connect(TEST_USER_ID).await; + + for downstream_uri in ["file:///password.env".to_owned(), format!("{}__file:///password.env", gateway.backend_name)] + { + let result = + service.read_resource(ReadResourceRequestParams::new(downstream_uri)).await.expect("resource is returned"); + + let Some(ResourceContents::TextResourceContents { text, uri, .. }) = result.contents.first() else { + panic!("expected text resource contents"); + }; + assert_eq!("file:///password.env", uri); + assert_ne!(BACKEND_RESOURCE_SECRET, text); + assert!(text.contains("[redacted]")); + } +} + #[tokio::test(flavor = "multi_thread", worker_threads = 1)] async fn pre_hook_rewrites_payload_without_changing_forwarded_parameter_headers() { let plugin = Arc::new(TestPlugin::new("pre", vec![cmf_hook_names::TOOL_PRE_INVOKE]).with_pre_rewrite()); diff --git a/crates/contextforge-data-plane-lib/tests/support/mod.rs b/crates/contextforge-data-plane-lib/tests/support/mod.rs index e40c84e6..94d9846e 100644 --- a/crates/contextforge-data-plane-lib/tests/support/mod.rs +++ b/crates/contextforge-data-plane-lib/tests/support/mod.rs @@ -25,8 +25,8 @@ pub(crate) use plugin::{ REWRITTEN_SUM_B, TestPlugin, TestPluginFactory, }; pub(crate) use plugin_gateway::{ - BACKEND_PROMPT_IMAGE, BACKEND_PROMPT_RESOURCE, RunningGateway, start_gateway, start_gateway_with_events, - start_gateway_with_json_backend_responses, start_gateway_with_parameter_headers, + BACKEND_PROMPT_IMAGE, BACKEND_PROMPT_RESOURCE, BACKEND_RESOURCE_SECRET, RunningGateway, start_gateway, + start_gateway_with_events, start_gateway_with_json_backend_responses, start_gateway_with_parameter_headers, }; pub(crate) use runtime::{runtime_with_post, runtime_with_pre, runtime_with_pre_and_post, runtime_with_prompt_plugin}; pub(crate) use test_gateways::{ diff --git a/crates/contextforge-data-plane-lib/tests/support/plugin_gateway.rs b/crates/contextforge-data-plane-lib/tests/support/plugin_gateway.rs index 65667e9f..cb4196ae 100644 --- a/crates/contextforge-data-plane-lib/tests/support/plugin_gateway.rs +++ b/crates/contextforge-data-plane-lib/tests/support/plugin_gateway.rs @@ -17,7 +17,8 @@ use rmcp::{ model::{ CallToolRequestParams, CallToolResponse, CallToolResult, ContentBlock, ErrorCode, GetPromptRequestParams, GetPromptResponse, GetPromptResult, Implementation, InitializeRequestParams, InitializeResult, NumberOrString, - ProgressNotificationParam, ProgressToken, PromptMessage, ResourceContents, Role, ServerCapabilities, + ProgressNotificationParam, ProgressToken, PromptMessage, ReadResourceRequestParams, ReadResourceResponse, + ReadResourceResult, ResourceContents, Role, ServerCapabilities, }, service::{RequestContext, Service}, transport::{ @@ -33,6 +34,7 @@ use super::{MemoryUserConfigStore, token}; pub(crate) const BACKEND_PROMPT_RESOURCE: &str = "token=secret"; pub(crate) const BACKEND_PROMPT_IMAGE: &str = "aW1hZ2UtYnl0ZXM="; +pub(crate) const BACKEND_RESOURCE_SECRET: &str = "AWS_ACCESS_KEY_ID=AKIAFAKE12345EXAMPLE"; // pragma: allowlist secret static GATEWAY_PORT_LOCK: OnceLock>> = OnceLock::new(); const CLIENT_CONNECT_TIMEOUT: Duration = Duration::from_secs(2); @@ -92,8 +94,18 @@ impl ServerHandler for TestBackend { _request: InitializeRequestParams, _cx: RequestContext, ) -> Result { - Ok(InitializeResult::new(ServerCapabilities::builder().enable_tools().enable_prompts().build()) - .with_server_info(Implementation::new("test-backend", "0.1.0"))) + Ok(InitializeResult::new( + ServerCapabilities::builder().enable_tools().enable_prompts().enable_resources().build(), + ) + .with_server_info(Implementation::new("test-backend", "0.1.0"))) + } + + async fn read_resource( + &self, + request: ReadResourceRequestParams, + _cx: RequestContext, + ) -> Result { + Ok(ReadResourceResult::new(vec![ResourceContents::text(BACKEND_RESOURCE_SECRET, request.uri)]).into()) } async fn get_prompt( @@ -237,7 +249,7 @@ pub const TOOL_NAMES: &[&str] = &[ "reflect_text", "wait_for_cancellation", ]; -pub const RESOURCE_URIS: &[&str] = &[""]; +pub const RESOURCE_URIS: &[&str] = &["file:///password.env"]; pub const PROMPT_NAMES: &[&str] = &["review_bundle", "review"]; pub(crate) struct RunningGateway { @@ -404,7 +416,12 @@ async fn start_gateway_with_state( .collect(), resource_uri_aliases: RESOURCE_URIS .iter() - .map(|n| NameAlias::new(n.to_string(), n.to_string())) + .flat_map(|uri| { + [ + NameAlias::new(uri.to_string(), uri.to_string()), + NameAlias::new(format!("{backend_name}__{uri}"), uri.to_string()), + ] + }) .collect(), prompt_name_aliases: PROMPT_NAMES .iter() diff --git a/crates/plugins/cpex-secrets-detection/README.md b/crates/plugins/cpex-secrets-detection/README.md index d77b96f2..d30c709d 100644 --- a/crates/plugins/cpex-secrets-detection/README.md +++ b/crates/plugins/cpex-secrets-detection/README.md @@ -21,7 +21,7 @@ Example config: { "name": "secrets-detection", "kind": "validator/secrets-detection", - "hooks": ["cmf.tool_pre_invoke", "cmf.tool_post_invoke"], + "hooks": ["cmf.tool_pre_invoke", "cmf.tool_post_invoke", "cmf.resource_post_fetch"], "config": { "redact": true, "redaction_text": "[redacted]", @@ -33,14 +33,12 @@ Example config: } ``` -The dataplane integration currently wires the tool-call path: +The dataplane integration supports: - `cmf.tool_pre_invoke`: scans tool arguments before the backend receives them. - `cmf.tool_post_invoke`: scans tool results before the client receives them. - -The crate also keeps prompt/resource stage handling for CPEX parity and future -hosts, but the current dataplane runtime config only uses the tool pre/post -hooks. +- `cmf.prompt_pre_fetch`: scans prompt arguments before rendering. +- `cmf.resource_post_fetch`: scans text resource contents before returning them. ## Behavior diff --git a/crates/plugins/cpex-secrets-detection/src/lib.rs b/crates/plugins/cpex-secrets-detection/src/lib.rs index 6d3259be..32abf830 100644 --- a/crates/plugins/cpex-secrets-detection/src/lib.rs +++ b/crates/plugins/cpex-secrets-detection/src/lib.rs @@ -54,8 +54,7 @@ impl Plugin for SecretsDetectionCore { #[derive(Debug, Clone, Copy, PartialEq, Eq)] pub enum Stage { - // The dataplane runtime config currently uses the tool pre/post stages. - // Prompt/resource stages are kept for CPEX parity and future hosts. + // Runtime configuration registers each stage independently. PromptPreFetch, ToolPreInvoke, ToolPostInvoke, From 592f6dce17ff6dd392af65b91a3f19644ddfa5ac Mon Sep 17 00:00:00 2001 From: lucarlig Date: Thu, 3 Sep 2026 11:20:03 +0100 Subject: [PATCH 2/3] refactor: centralize CPEX hook plumbing Signed-off-by: lucarlig --- .../contextforge-data-plane-cpex/src/cmf.rs | 59 ++++------- .../src/factory.rs | 4 +- .../src/runtime.rs | 99 ++++--------------- 3 files changed, 40 insertions(+), 122 deletions(-) diff --git a/crates/contextforge-data-plane-cpex/src/cmf.rs b/crates/contextforge-data-plane-cpex/src/cmf.rs index 5f1a66a8..06a5229c 100644 --- a/crates/contextforge-data-plane-cpex/src/cmf.rs +++ b/crates/contextforge-data-plane-cpex/src/cmf.rs @@ -80,23 +80,7 @@ pub(crate) fn resource_result_payload( .contents .iter() .map(|content| { - let (uri, mime_type, text) = match content { - 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, - }; - Some(ContentPart::Resource { - content: CmfResource { - resource_request_id: resource_request_id.to_owned(), - uri, - resource_type: ResourceType::Uri, - content: text, - mime_type, - ..Default::default() - }, - }) + cmf_resource_content(content, resource_request_id).map(|content| ContentPart::Resource { content }) }) .collect::>>()?; Some(MessagePayload { @@ -104,6 +88,24 @@ pub(crate) fn resource_result_payload( }) } +fn cmf_resource_content(content: &ResourceContents, resource_request_id: &str) -> Option { + let (uri, mime_type, text) = match content { + 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, + }; + Some(CmfResource { + resource_request_id: resource_request_id.to_owned(), + uri, + resource_type: ResourceType::Uri, + content: text, + mime_type, + ..Default::default() + }) +} + pub(crate) fn resource_result_response( mut original: ReadResourceResult, payload: &MessagePayload, @@ -379,28 +381,7 @@ fn cmf_content_part(block: &ContentBlock, prompt_request_id: &str) -> Option { - 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, - }, - } + ContentPart::Resource { content: cmf_resource_content(&resource.resource, prompt_request_id)? } }, ContentBlock::ResourceLink(link) => ContentPart::ResourceRef { content: ResourceReference { diff --git a/crates/contextforge-data-plane-cpex/src/factory.rs b/crates/contextforge-data-plane-cpex/src/factory.rs index a2c8c6e3..627abe1c 100644 --- a/crates/contextforge-data-plane-cpex/src/factory.rs +++ b/crates/contextforge-data-plane-cpex/src/factory.rs @@ -28,7 +28,7 @@ where let handlers = config .hooks .iter() - .filter_map(|hook| cmf_hook_name(hook)) + .filter_map(|hook| supported_cmf_hook_name(hook)) .map(|hook| { (hook, Arc::new(TypedHandlerAdapter::::new(Arc::clone(&plugin))) as Arc) }) @@ -45,7 +45,7 @@ where } } -fn cmf_hook_name(hook: &str) -> Option<&'static str> { +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), diff --git a/crates/contextforge-data-plane-cpex/src/runtime.rs b/crates/contextforge-data-plane-cpex/src/runtime.rs index fd9d0ce6..11243e10 100644 --- a/crates/contextforge-data-plane-cpex/src/runtime.rs +++ b/crates/contextforge-data-plane-cpex/src/runtime.rs @@ -25,6 +25,7 @@ use crate::{ tool_call_payload, tool_json_result_payload, tool_result_payload, }, error::GatewayPluginRuntimeError, + factory::supported_cmf_hook_name, hooks::{PromptPreFetchResult, ResourcePreFetchResult, RuntimeHookState, ToolArgumentsUpdate, ToolPreCallResult}, pipeline::{ effective_post_json, effective_post_prompt_result, effective_post_resource_result, effective_post_result, @@ -152,15 +153,6 @@ impl Drop for GatewayPluginRuntime { } } -const SUPPORTED_HOOKS: [&str; 6] = [ - cmf_hook_names::TOOL_PRE_INVOKE, - cmf_hook_names::TOOL_POST_INVOKE, - cmf_hook_names::PROMPT_PRE_FETCH, - cmf_hook_names::PROMPT_POST_FETCH, - cmf_hook_names::RESOURCE_PRE_FETCH, - cmf_hook_names::RESOURCE_POST_FETCH, -]; - fn declares(config: &CpexConfig, hook_name: &str) -> bool { config.plugins.iter().any(|plugin| plugin.hooks.iter().any(|hook| hook == hook_name)) } @@ -181,7 +173,7 @@ fn validate_gateway_supported_config(config: &CpexConfig) -> Result<(), GatewayP return Err(GatewayPluginRuntimeError::ConfigUnsupported); } - if plugin.hooks.iter().any(|hook| !SUPPORTED_HOOKS.contains(&hook.as_str())) { + if plugin.hooks.iter().any(|hook| supported_cmf_hook_name(hook).is_none()) { return Err(GatewayPluginRuntimeError::ConfigUnsupported); } } @@ -190,26 +182,15 @@ fn validate_gateway_supported_config(config: &CpexConfig) -> Result<(), GatewayP } impl GatewayPluginRuntime { - async fn invoke_tool_pre(&self, payload: MessagePayload) -> PipelineResult { - let (result, background_tasks) = self - .manager - .invoke_named::(cmf_hook_names::TOOL_PRE_INVOKE, payload, Extensions::default(), None) - .await; - log_pipeline_errors(cmf_hook_names::TOOL_PRE_INVOKE, &result); - drop(background_tasks); - result - } - - async fn invoke_tool_post( + async fn invoke_cmf_hook( &self, + hook_name: &'static str, payload: MessagePayload, context_table: Option, ) -> PipelineResult { - let (result, background_tasks) = self - .manager - .invoke_named::(cmf_hook_names::TOOL_POST_INVOKE, payload, Extensions::default(), context_table) - .await; - log_pipeline_errors(cmf_hook_names::TOOL_POST_INVOKE, &result); + 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 } @@ -227,7 +208,7 @@ impl GatewayPluginRuntime { 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_tool_pre(original_payload).await; + 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)); } @@ -237,30 +218,6 @@ impl GatewayPluginRuntime { Ok(ToolPreCallResult { arguments, state: Some(Arc::new(state)) }) } - async fn invoke_prompt_pre(&self, payload: MessagePayload) -> PipelineResult { - let (result, background_tasks) = self - .manager - .invoke_named::(cmf_hook_names::PROMPT_PRE_FETCH, payload, Extensions::default(), None) - .await; - log_pipeline_errors(cmf_hook_names::PROMPT_PRE_FETCH, &result); - drop(background_tasks); - result - } - - async fn invoke_prompt_post( - &self, - payload: MessagePayload, - context_table: Option, - ) -> PipelineResult { - let (result, background_tasks) = self - .manager - .invoke_named::(cmf_hook_names::PROMPT_POST_FETCH, payload, Extensions::default(), context_table) - .await; - log_pipeline_errors(cmf_hook_names::PROMPT_POST_FETCH, &result); - drop(background_tasks); - result - } - pub(crate) async fn before_get_prompt( &self, request: &GetPromptRequestParams, @@ -279,7 +236,7 @@ impl GatewayPluginRuntime { let prompt_request_id = next_prompt_request_id(); let payload = prompt_request_payload(request, prompt_name, backend_name, &prompt_request_id); - let pre_result = self.invoke_prompt_pre(payload).await; + 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)); } @@ -296,30 +253,6 @@ impl GatewayPluginRuntime { Ok(PromptPreFetchResult { arguments, state }) } - async fn invoke_resource_pre(&self, payload: MessagePayload) -> PipelineResult { - let (result, background_tasks) = self - .manager - .invoke_named::(cmf_hook_names::RESOURCE_PRE_FETCH, payload, Extensions::default(), None) - .await; - log_pipeline_errors(cmf_hook_names::RESOURCE_PRE_FETCH, &result); - drop(background_tasks); - result - } - - async fn invoke_resource_post( - &self, - payload: MessagePayload, - context_table: Option, - ) -> PipelineResult { - let (result, background_tasks) = self - .manager - .invoke_named::(cmf_hook_names::RESOURCE_POST_FETCH, payload, Extensions::default(), context_table) - .await; - log_pipeline_errors(cmf_hook_names::RESOURCE_POST_FETCH, &result); - drop(background_tasks); - result - } - pub(crate) async fn before_read_resource(&self, resource_uri: &str) -> Result { let resource_request_id = next_resource_request_id(); if !self.hooks.resource.pre { @@ -328,7 +261,7 @@ impl GatewayPluginRuntime { } let payload = resource_request_payload(resource_uri, &resource_request_id); - let pre_result = self.invoke_resource_pre(payload).await; + 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)); } @@ -355,7 +288,8 @@ impl GatewayPluginRuntime { let Some(state) = state else { return Ok(response) }; let payload = prompt_result_payload(&response, prompt_name, &state.prompt_request_id); - let post_result = self.invoke_prompt_post(payload, Some(state.context_table.clone())).await; + 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)); } @@ -376,7 +310,8 @@ impl GatewayPluginRuntime { let Some(state) = state else { return Ok(response) }; 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_resource_post(payload, state.context_table.clone()).await; + let post_result = + self.invoke_cmf_hook(cmf_hook_names::RESOURCE_POST_FETCH, payload, state.context_table.clone()).await; if post_result.is_denied() { return Err(plugin_denied_error("resource", post_result)); } @@ -398,7 +333,8 @@ impl GatewayPluginRuntime { let mut state = state.lock().await; let post_result = self - .invoke_tool_post( + .invoke_cmf_hook( + cmf_hook_names::TOOL_POST_INVOKE, tool_result_payload(tool_name, &response, &state.tool_call_id), Some(state.context_table.clone()), ) @@ -430,7 +366,8 @@ impl GatewayPluginRuntime { let content = serde_json::to_value(&event).unwrap_or(serde_json::Value::Null); let mut state = state.lock().await; let post_result = self - .invoke_tool_post( + .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()), ) From 6df91b2d77f11221e3733fa14036024d890b6c64 Mon Sep 17 00:00:00 2001 From: lucarlig Date: Thu, 3 Sep 2026 11:59:04 +0100 Subject: [PATCH 3/3] fix: harden CPEX resource hook lifecycle Signed-off-by: lucarlig --- Cargo.lock | 1 + Cargo.toml | 1 + _context/wiki/config.md | 4 +- .../contextforge-data-plane-cpex/Cargo.toml | 1 + .../contextforge-data-plane-cpex/src/cmf.rs | 171 +++++++++++++++--- .../src/handle.rs | 104 +++++++++-- .../contextforge-data-plane-cpex/src/hooks.rs | 30 ++- .../contextforge-data-plane-cpex/src/lib.rs | 4 +- .../src/runtime.rs | 44 +++-- crates/contextforge-data-plane-lib/Cargo.toml | 2 +- 10 files changed, 299 insertions(+), 63 deletions(-) diff --git a/Cargo.lock b/Cargo.lock index 64485c8b..6168751f 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -582,6 +582,7 @@ version = "0.1.0" dependencies = [ "arc-swap", "async-trait", + "base64 0.22.1", "contextforge-data-plane-apis", "cpex", "redis", diff --git a/Cargo.toml b/Cargo.toml index 6a211c8e..4e474b9e 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -51,6 +51,7 @@ clap = { version = "4.5.60", features = ["derive", "env"] } thiserror = "2.0.18" rmp-serde = "1.3.1" async-trait = "0.1.89" +base64 = "0.22.1" reqwest = "0.13" jsonwebtoken = { version = "11.0.0", features = ["aws_lc_rs"] } rustls = { version = "0.23", features = ["ring"] } diff --git a/_context/wiki/config.md b/_context/wiki/config.md index 6059df2b..0493aef9 100644 --- a/_context/wiki/config.md +++ b/_context/wiki/config.md @@ -201,11 +201,11 @@ Writing plugin edits back follows three rules: MCP prompt results carry no error flag, so a plugin setting `is_error` on the CMF prompt result is rejecting the prompt rather than describing it. The gateway turns that into an MCP error carrying the plugin's `error_message`, and the rendered content never reaches the client. This differs from tools, where `is_error` is a field on `CallToolResult` and is forwarded as a successful response. -Binary resource blobs reach plugins by URI and MIME type but not by content: CMF stores decoded bytes while MCP sends base64. A plugin can deny such a message; editing one fails the write-back. +Binary resource blobs reach plugins as decoded CMF bytes and are encoded back to MCP base64 after an edit. Unchanged blobs retain the backend's exact wire representation. ### Resource Read Hook Behavior -For `resources/read`, the pre hook runs after routing and may allow or deny the canonical backend-local URI, but cannot change it. The post hook can redact text resources or deny the response. URI, MIME type, item count, and blob identity must remain stable; unsupported or lossy edits fail closed. +For `resources/read`, the pre hook runs after routing and may allow or deny the canonical backend-local URI, but cannot change it. The post hook can redact text resources, transform binary resources, or deny the response. URI, MIME type, item count, CMF schema version, and channel must remain stable; unsupported or lossy edits and invalid hook lifecycle state fail closed. ### Demo Plugin Workflow diff --git a/crates/contextforge-data-plane-cpex/Cargo.toml b/crates/contextforge-data-plane-cpex/Cargo.toml index d74d69bd..6df8e734 100644 --- a/crates/contextforge-data-plane-cpex/Cargo.toml +++ b/crates/contextforge-data-plane-cpex/Cargo.toml @@ -15,6 +15,7 @@ doctest = false [dependencies] arc-swap = "1.7" async-trait.workspace = true +base64.workspace = true contextforge-data-plane-apis.workspace = true cpex.workspace = true redis.workspace = true diff --git a/crates/contextforge-data-plane-cpex/src/cmf.rs b/crates/contextforge-data-plane-cpex/src/cmf.rs index 06a5229c..eab86bfc 100644 --- a/crates/contextforge-data-plane-cpex/src/cmf.rs +++ b/crates/contextforge-data-plane-cpex/src/cmf.rs @@ -1,8 +1,9 @@ use std::collections::HashMap; +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, + Resource as CmfResource, ResourceReference, ResourceType, Role, ToolCall, ToolResult, constants::SCHEMA_VERSION, }; use rmcp::model::{ CallToolRequestParams, CallToolResult, ContentBlock, GetPromptRequestParams, GetPromptResult, PromptMessage, @@ -18,7 +19,7 @@ pub(crate) fn tool_call_payload( ) -> MessagePayload { MessagePayload { message: Message { - schema_version: "2.0".to_owned(), + schema_version: SCHEMA_VERSION.to_owned(), role: Role::Assistant, content: vec![ContentPart::ToolCall { content: ToolCall { @@ -36,7 +37,7 @@ pub(crate) fn tool_call_payload( pub(crate) fn resource_request_payload(resource_uri: &str, resource_request_id: &str) -> MessagePayload { MessagePayload { message: Message { - schema_version: "2.0".to_owned(), + schema_version: SCHEMA_VERSION.to_owned(), role: Role::User, content: vec![ContentPart::ResourceRef { content: ResourceReference { @@ -62,7 +63,7 @@ pub(crate) fn resource_request_matches( let [ContentPart::ResourceRef { content }] = payload.message.content.as_slice() else { return false; }; - payload.message.role == Role::User + canonical_message_envelope(payload, Role::User) && content.resource_request_id == resource_request_id && content.uri == resource_uri && matches!(content.resource_type, ResourceType::Uri) @@ -84,16 +85,18 @@ pub(crate) fn resource_result_payload( }) .collect::>>()?; Some(MessagePayload { - message: Message { schema_version: "2.0".to_owned(), role: Role::Assistant, content, channel: None }, + 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) = match content { + let (uri, mime_type, text, blob) = match content { ResourceContents::TextResourceContents { uri, mime_type, text, .. } => { - (uri.clone(), mime_type.clone(), Some(text.clone())) + (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()?)) }, - ResourceContents::BlobResourceContents { uri, mime_type, .. } => (uri.clone(), mime_type.clone(), None), _ => return None, }; Some(CmfResource { @@ -101,6 +104,7 @@ fn cmf_resource_content(content: &ResourceContents, resource_request_id: &str) - uri, resource_type: ResourceType::Uri, content: text, + blob, mime_type, ..Default::default() }) @@ -111,7 +115,8 @@ pub(crate) fn resource_result_response( payload: &MessagePayload, resource_request_id: &str, ) -> Option { - if payload.message.role != Role::Assistant || payload.message.content.len() != original.contents.len() { + if !canonical_message_envelope(payload, Role::Assistant) || payload.message.content.len() != original.contents.len() + { return None; } @@ -123,7 +128,6 @@ pub(crate) fn resource_result_response( || !matches!(content.resource_type, ResourceType::Uri) || content.name.is_some() || content.description.is_some() - || content.blob.is_some() || content.size_bytes.is_some() || !content.annotations.is_empty() || content.version.is_some() @@ -132,15 +136,20 @@ pub(crate) fn resource_result_response( } match original { ResourceContents::TextResourceContents { uri, mime_type, text, .. } => { - if content.uri != *uri || content.mime_type != *mime_type { + if content.uri != *uri || content.mime_type != *mime_type || content.blob.is_some() { return None; } *text = content.content.clone()?; }, - ResourceContents::BlobResourceContents { uri, mime_type, .. } => { + ResourceContents::BlobResourceContents { uri, mime_type, blob, .. } => { if content.uri != *uri || content.mime_type != *mime_type || content.content.is_some() { return None; } + let modified_blob = content.blob.as_ref()?; + let original_blob = BASE64_STANDARD.decode(blob.as_bytes()).ok()?; + if modified_blob != &original_blob { + *blob = BASE64_STANDARD.encode(modified_blob); + } }, _ => return None, } @@ -148,6 +157,12 @@ pub(crate) fn resource_result_response( Some(original) } +fn canonical_message_envelope(payload: &MessagePayload, role: Role) -> bool { + payload.message.schema_version == SCHEMA_VERSION + && payload.message.role == role + && payload.message.channel.is_none() +} + pub(crate) fn tool_result_payload(tool_name: &str, response: &CallToolResult, tool_call_id: &str) -> MessagePayload { tool_json_result_payload( tool_name, @@ -165,7 +180,7 @@ pub(crate) fn tool_json_result_payload( ) -> MessagePayload { MessagePayload { message: Message { - schema_version: "2.0".to_owned(), + schema_version: SCHEMA_VERSION.to_owned(), role: Role::Tool, content: vec![ContentPart::ToolResult { content: ToolResult { @@ -241,7 +256,7 @@ pub(crate) fn prompt_request_payload( ) -> MessagePayload { MessagePayload { message: Message { - schema_version: "2.0".to_owned(), + schema_version: SCHEMA_VERSION.to_owned(), role: Role::User, content: vec![ContentPart::PromptRequest { content: PromptRequest { @@ -283,7 +298,7 @@ pub(crate) fn prompt_result_payload( MessagePayload { message: Message { - schema_version: "2.0".to_owned(), + schema_version: SCHEMA_VERSION.to_owned(), role: Role::Assistant, content: vec![ContentPart::PromptResult { content: PromptResult { @@ -352,7 +367,7 @@ pub(crate) fn prompt_result_response( fn cmf_prompt_message(message: &PromptMessage, prompt_request_id: &str) -> Message { Message { - schema_version: "2.0".to_owned(), + schema_version: SCHEMA_VERSION.to_owned(), role: match message.role { McpRole::Assistant => Role::Assistant, McpRole::User => Role::User, @@ -422,12 +437,7 @@ fn mcp_prompt_message(message: &Message) -> Option { 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::Resource { content } => ContentBlock::resource(mcp_resource_content(content)?), ContentPart::ResourceRef { content } => { ContentBlock::ResourceLink(McpResource::new(content.uri.clone(), content.name.clone()?)) }, @@ -437,8 +447,28 @@ fn mcp_prompt_message(message: &Message) -> Option { Some(PromptMessage::new(role, content)) } +fn mcp_resource_content(content: &CmfResource) -> Option { + match (&content.content, &content.blob) { + (Some(text), None) => Some(ResourceContents::TextResourceContents { + uri: content.uri.clone(), + mime_type: content.mime_type.clone(), + text: text.clone(), + meta: None, + }), + (None, Some(blob)) => Some(ResourceContents::BlobResourceContents { + uri: content.uri.clone(), + mime_type: content.mime_type.clone(), + blob: BASE64_STANDARD.encode(blob), + meta: None, + }), + _ => None, + } +} + #[cfg(test)] mod tests { + use cpex::cpex_core::cmf::Channel; + use super::*; #[test] @@ -452,6 +482,19 @@ mod tests { assert!(!resource_request_matches(&payload, "file:///password.env", "resource-1")); } + #[test] + fn resource_request_rejects_modified_envelope() { + let mut payload = resource_request_payload("file:///password.env", "resource-1"); + payload.message.schema_version = "3.0".to_owned(); + + assert!(!resource_request_matches(&payload, "file:///password.env", "resource-1")); + + let mut payload = resource_request_payload("file:///password.env", "resource-1"); + payload.message.channel = Some(Channel::Analysis); + + assert!(!resource_request_matches(&payload, "file:///password.env", "resource-1")); + } + #[test] fn resource_result_response_applies_only_text_changes() { let original = @@ -483,6 +526,65 @@ mod tests { assert!(resource_result_response(original, &payload, "resource-1").is_none()); } + #[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, "resource-1").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, "resource-1").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_response_rejects_modified_envelope() { + let original = ReadResourceResult::new(vec![ResourceContents::text("secret", "file:///password.env")]); + let mut payload = resource_result_payload(&original, "resource-1").expect("resource result is supported"); + payload.message.schema_version = "3.0".to_owned(); + + assert!(resource_result_response(original.clone(), &payload, "resource-1").is_none()); + + let mut payload = resource_result_payload(&original, "resource-1").expect("resource result is supported"); + payload.message.channel = Some(Channel::Final); + + assert!(resource_result_response(original, &payload, "resource-1").is_none()); + } + fn text_prompt() -> GetPromptResult { GetPromptResult::new(vec![PromptMessage::new_text(McpRole::User, "review of weather")]) } @@ -918,6 +1020,31 @@ mod tests { assert_eq!("file:///app.env", uri); } + #[test] + fn prompt_result_response_round_trips_embedded_blob_resource() { + let original = GetPromptResult::new(vec![PromptMessage::new( + McpRole::User, + ContentBlock::resource(ResourceContents::blob(BASE64_STANDARD.encode(b"token=secret"), "file:///app.bin")), + )]); + 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(b"token=secret".as_slice()), content.blob.as_deref()); + content.blob = Some(b"token=[REDACTED]".to_vec()); + + 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::BlobResourceContents { blob, uri, .. } = &resource.resource else { + panic!("expected blob resource contents"); + }; + assert_eq!(b"token=[REDACTED]", BASE64_STANDARD.decode(blob).expect("valid base64").as_slice()); + assert_eq!("file:///app.bin", uri); + } + #[test] fn tool_result_response_uses_cmf_error_flag_for_nested_mcp_result() { let original = CallToolResult::success(vec![ContentBlock::text("original")]); diff --git a/crates/contextforge-data-plane-cpex/src/handle.rs b/crates/contextforge-data-plane-cpex/src/handle.rs index 6dcd0dd8..5fd30c31 100644 --- a/crates/contextforge-data-plane-cpex/src/handle.rs +++ b/crates/contextforge-data-plane-cpex/src/handle.rs @@ -23,7 +23,10 @@ use tokio::task::JoinHandle; use crate::{ config::{LoadedRuntimePluginConfig, RedisRuntimePluginConfigStore, RuntimePluginConfigStore, cpex_config}, error::GatewayPluginRuntimeError, - hooks::{PromptPreFetchResult, ResourcePreFetchResult, RuntimeHookError, RuntimeHookState, ToolPreCallResult}, + hooks::{ + PromptPreFetchResult, ResourceHookState, ResourcePreFetchResult, RuntimeHookError, RuntimeHookState, + ToolPreCallResult, invalid_resource_hook_state_error, + }, runtime::GatewayPluginRuntime, }; @@ -47,6 +50,11 @@ struct RegistryCallState { state: Option, } +struct RegistryResourceCallState { + runtime: Arc, + state: RuntimeHookState, +} + enum RuntimeState { Active(Arc), Failed(String), @@ -287,14 +295,13 @@ impl GatewayPluginRuntimeHandle { let RuntimeState::Active(runtime) = state.as_ref() else { return Err(runtime_failed_error(state.as_ref())); }; - let mut result = runtime.before_read_resource(resource_uri).await?; - if runtime.has_resource_post_hook() { - let state = result.state.take(); - result.state = Some(Arc::new(RegistryCallState { runtime: Arc::clone(runtime), state })); - } else { - result.state = None; + match runtime.before_read_resource(resource_uri).await?.state.into_inner() { + Some(state) => Ok(ResourcePreFetchResult::with_post_state(Arc::new(RegistryResourceCallState { + runtime: Arc::clone(runtime), + state, + }))), + None => Ok(ResourcePreFetchResult::unchanged()), } - Ok(result) } pub async fn after_get_prompt( @@ -312,11 +319,19 @@ impl GatewayPluginRuntimeHandle { pub async fn after_read_resource( &self, response: ReadResourceResult, - state: Option, + state: ResourceHookState, ) -> Result { - match state.and_then(|state| state.downcast::().ok()) { - Some(state) => state.runtime.after_read_resource(response, state.state.clone()).await, - None => Ok(response), + let Some(state) = state.into_inner() else { + let current = self.current(); + return match current.as_ref() { + RuntimeState::Active(runtime) if !runtime.has_resource_post_hook() => Ok(response), + RuntimeState::Active(_) => Err(invalid_resource_hook_state_error()), + failed @ RuntimeState::Failed(_) => Err(runtime_failed_error(failed)), + }; + }; + match state.downcast::() { + Ok(state) => state.runtime.after_read_resource(response, Arc::clone(&state.state)).await, + Err(_) => Err(invalid_resource_hook_state_error()), } } @@ -370,7 +385,7 @@ mod tests { use async_trait::async_trait; use cpex::cpex_core::{ - cmf::{CmfHook, ContentPart, MessagePayload, Role}, + cmf::{CmfHook, ContentPart, MessagePayload}, context::PluginContext, error::{PluginError, PluginViolation}, factory::{PluginFactory, PluginInstance}, @@ -380,6 +395,7 @@ mod tests { }; use rmcp::model::{ CallToolRequestParams, CallToolResult, ContentBlock, NumberOrString, ProgressNotificationParam, ProgressToken, + ReadResourceResult, ResourceContents, }; use serde_json::{Value, json}; use tokio::sync::Mutex as TokioMutex; @@ -545,7 +561,12 @@ mod tests { _extensions: &Extensions, ctx: &mut PluginContext, ) -> PluginResult { - let is_post = payload.message.role == Role::Tool; + 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; @@ -669,6 +690,7 @@ mod tests { 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(( @@ -809,6 +831,60 @@ mod tests { assert_eq!(1, observations.lock().expect("observations lock poisoned").pre_calls); } + #[tokio::test(flavor = "multi_thread", worker_threads = 1)] + async fn resource_post_hook_rejects_missing_lifecycle_state() { + let plugin = Arc::new(TestPlugin::new("resource", vec![cmf_hook_names::RESOURCE_POST_FETCH])); + let runtime = runtime_with_plugin(&plugin, plugin_config(&[Arc::clone(&plugin)])).await; + let response = ReadResourceResult::new(vec![ResourceContents::text("secret", "file:///password.env")]); + + let error = runtime + .handle() + .after_read_resource(response, ResourcePreFetchResult::unchanged().state) + .await + .expect_err("missing state fails closed"); + + assert_eq!(ErrorCode::INTERNAL_ERROR, error.code); + assert_eq!("Resource post-hook state is missing or invalid", error.message); + } + + #[tokio::test(flavor = "multi_thread", worker_threads = 1)] + async fn resource_post_hook_rejects_invalid_lifecycle_state() { + let plugin = Arc::new(TestPlugin::new("resource", vec![cmf_hook_names::RESOURCE_POST_FETCH])); + let runtime = runtime_with_plugin(&plugin, plugin_config(&[Arc::clone(&plugin)])).await; + let response = ReadResourceResult::new(vec![ResourceContents::text("secret", "file:///password.env")]); + + let error = runtime + .handle() + .after_read_resource(response, ResourceHookState::active(Arc::new(()))) + .await + .expect_err("invalid state fails closed"); + + assert_eq!(ErrorCode::INTERNAL_ERROR, error.code); + assert_eq!("Resource post-hook state is missing or invalid", 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")]); + + runtime + .handle() + .after_read_resource(response, pre.state) + .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 = diff --git a/crates/contextforge-data-plane-cpex/src/hooks.rs b/crates/contextforge-data-plane-cpex/src/hooks.rs index 0b481601..392d0fd0 100644 --- a/crates/contextforge-data-plane-cpex/src/hooks.rs +++ b/crates/contextforge-data-plane-cpex/src/hooks.rs @@ -1,6 +1,9 @@ use std::{any::Any, sync::Arc}; -use rmcp::model::{CallToolRequestParams, GetPromptRequestParams}; +use rmcp::{ + ErrorData, + model::{CallToolRequestParams, GetPromptRequestParams}, +}; use serde_json::{Map, Value}; pub type RuntimeHookError = Box; @@ -53,15 +56,36 @@ pub struct PromptPreFetchResult { } pub struct ResourcePreFetchResult { - pub state: Option, + pub state: ResourceHookState, +} + +/// Opaque state connecting a resource pre-fetch hook to its post-fetch hook. +pub struct ResourceHookState(Option); + +impl ResourceHookState { + pub(crate) fn active(state: RuntimeHookState) -> Self { + Self(Some(state)) + } + + pub(crate) fn into_inner(self) -> Option { + self.0 + } } impl ResourcePreFetchResult { pub fn unchanged() -> Self { - Self { state: None } + Self { state: ResourceHookState(None) } + } + + pub(crate) fn with_post_state(state: RuntimeHookState) -> Self { + Self { state: ResourceHookState::active(state) } } } +pub(crate) fn invalid_resource_hook_state_error() -> ErrorData { + ErrorData::internal_error("Resource post-hook state is missing or invalid", None) +} + impl PromptPreFetchResult { pub fn unchanged() -> Self { Self { arguments: PromptArgumentsUpdate::Unchanged, state: None } diff --git a/crates/contextforge-data-plane-cpex/src/lib.rs b/crates/contextforge-data-plane-cpex/src/lib.rs index 944bfb22..e51bd818 100644 --- a/crates/contextforge-data-plane-cpex/src/lib.rs +++ b/crates/contextforge-data-plane-cpex/src/lib.rs @@ -11,6 +11,6 @@ pub use error::GatewayPluginRuntimeError; pub use factory::CmfPluginFactory; pub use handle::{CpexRuntimeRegistry, GatewayPluginRuntimeHandle}; pub use hooks::{ - PromptArgumentsUpdate, PromptPreFetchResult, ResourcePreFetchResult, RuntimeHookError, RuntimeHookState, - ToolArgumentsUpdate, ToolPreCallResult, + PromptArgumentsUpdate, PromptPreFetchResult, ResourceHookState, ResourcePreFetchResult, RuntimeHookError, + RuntimeHookState, ToolArgumentsUpdate, ToolPreCallResult, }; diff --git a/crates/contextforge-data-plane-cpex/src/runtime.rs b/crates/contextforge-data-plane-cpex/src/runtime.rs index 11243e10..28c4a712 100644 --- a/crates/contextforge-data-plane-cpex/src/runtime.rs +++ b/crates/contextforge-data-plane-cpex/src/runtime.rs @@ -26,7 +26,10 @@ use crate::{ }, error::GatewayPluginRuntimeError, factory::supported_cmf_hook_name, - hooks::{PromptPreFetchResult, ResourcePreFetchResult, RuntimeHookState, ToolArgumentsUpdate, ToolPreCallResult}, + hooks::{ + PromptPreFetchResult, ResourcePreFetchResult, RuntimeHookState, ToolArgumentsUpdate, ToolPreCallResult, + invalid_resource_hook_state_error, + }, pipeline::{ effective_post_json, effective_post_prompt_result, effective_post_resource_result, effective_post_result, effective_pre_args, effective_pre_prompt_args, log_pipeline_errors, plugin_denied_error, @@ -87,7 +90,7 @@ fn new_prompt_call_state(context_table: PluginContextTable, prompt_request_id: S } struct ResourceCallState { - context_table: Option, + context_table: PluginContextTable, resource_request_id: String, } @@ -95,7 +98,7 @@ fn next_resource_request_id() -> String { format!("gateway-resource-request-{}", CORRELATION_ID.fetch_add(1, Ordering::Relaxed)) } -fn new_resource_call_state(context_table: Option, resource_request_id: String) -> RuntimeHookState { +fn new_resource_call_state(context_table: PluginContextTable, resource_request_id: String) -> RuntimeHookState { Arc::new(ResourceCallState { context_table, resource_request_id }) } @@ -254,10 +257,16 @@ impl GatewayPluginRuntime { } pub(crate) async fn before_read_resource(&self, resource_uri: &str) -> Result { + if !self.hooks.resource.pre && !self.hooks.resource.post { + return Ok(ResourcePreFetchResult::unchanged()); + } + let resource_request_id = next_resource_request_id(); if !self.hooks.resource.pre { - let state = self.hooks.resource.post.then(|| new_resource_call_state(None, resource_request_id)); - return Ok(ResourcePreFetchResult { state }); + return Ok(ResourcePreFetchResult::with_post_state(new_resource_call_state( + PluginContextTable::default(), + resource_request_id, + ))); } let payload = resource_request_payload(resource_uri, &resource_request_id); @@ -266,12 +275,14 @@ impl GatewayPluginRuntime { return Err(plugin_denied_error("resource", pre_result)); } validate_pre_resource_result(&pre_result, resource_uri, &resource_request_id)?; - let state = self - .hooks - .resource - .post - .then(|| new_resource_call_state(Some(pre_result.context_table), resource_request_id)); - Ok(ResourcePreFetchResult { state }) + if self.hooks.resource.post { + Ok(ResourcePreFetchResult::with_post_state(new_resource_call_state( + pre_result.context_table, + resource_request_id, + ))) + } else { + Ok(ResourcePreFetchResult::unchanged()) + } } pub(crate) async fn after_get_prompt( @@ -300,18 +311,13 @@ impl GatewayPluginRuntime { pub(crate) async fn after_read_resource( &self, response: ReadResourceResult, - state: Option, + state: RuntimeHookState, ) -> Result { - if !self.hooks.resource.post { - return Ok(response); - } - - let state = state.and_then(|state| state.downcast::().ok()); - let Some(state) = state else { return Ok(response) }; + let state = state.downcast::().map_err(|_| invalid_resource_hook_state_error())?; 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, state.context_table.clone()).await; + self.invoke_cmf_hook(cmf_hook_names::RESOURCE_POST_FETCH, payload, Some(state.context_table.clone())).await; if post_result.is_denied() { return Err(plugin_denied_error("resource", post_result)); } diff --git a/crates/contextforge-data-plane-lib/Cargo.toml b/crates/contextforge-data-plane-lib/Cargo.toml index 681dabab..427f91d5 100644 --- a/crates/contextforge-data-plane-lib/Cargo.toml +++ b/crates/contextforge-data-plane-lib/Cargo.toml @@ -35,7 +35,7 @@ clap.workspace = true thiserror.workspace = true rmp-serde.workspace = true async-trait.workspace = true -base64 = "0.22.1" +base64.workspace = true reqwest.workspace = true uuid.workspace = true lru_time_cache = "0.11.11"