fix(goose): propagate session_id across providers and MCP (#6584)
Signed-off-by: Adrian Cole <adrian@tetrate.io>
This commit is contained in:
@@ -26,47 +26,35 @@ use rmcp::{
|
||||
ClientHandler, ErrorData, Peer, RoleClient, ServiceError, ServiceExt,
|
||||
};
|
||||
use serde_json::Value;
|
||||
use std::{sync::Arc, time::Duration};
|
||||
use std::{
|
||||
sync::{Arc, OnceLock},
|
||||
time::Duration,
|
||||
};
|
||||
use tokio::sync::{
|
||||
mpsc::{self, Sender},
|
||||
Mutex,
|
||||
};
|
||||
use tokio_util::sync::CancellationToken;
|
||||
use uuid::Uuid;
|
||||
|
||||
pub type BoxError = Box<dyn std::error::Error + Sync + Send>;
|
||||
|
||||
pub type Error = rmcp::ServiceError;
|
||||
|
||||
#[derive(Clone, Debug)]
|
||||
pub struct McpMeta {
|
||||
pub session_id: String,
|
||||
}
|
||||
|
||||
impl McpMeta {
|
||||
pub fn new(session_id: impl Into<String>) -> Self {
|
||||
Self {
|
||||
session_id: session_id.into(),
|
||||
}
|
||||
}
|
||||
|
||||
fn inject_into_extensions(&self, extensions: Extensions) -> Extensions {
|
||||
inject_session_id_into_extensions(extensions, &self.session_id)
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait::async_trait]
|
||||
pub trait McpClientTrait: Send + Sync {
|
||||
async fn list_tools(
|
||||
&self,
|
||||
session_id: &str,
|
||||
next_cursor: Option<String>,
|
||||
cancel_token: CancellationToken,
|
||||
) -> Result<ListToolsResult, Error>;
|
||||
|
||||
async fn call_tool(
|
||||
&self,
|
||||
session_id: &str,
|
||||
name: &str,
|
||||
arguments: Option<JsonObject>,
|
||||
meta: McpMeta,
|
||||
cancel_token: CancellationToken,
|
||||
) -> Result<CallToolResult, Error>;
|
||||
|
||||
@@ -74,6 +62,7 @@ pub trait McpClientTrait: Send + Sync {
|
||||
|
||||
async fn list_resources(
|
||||
&self,
|
||||
_session_id: &str,
|
||||
_next_cursor: Option<String>,
|
||||
_cancel_token: CancellationToken,
|
||||
) -> Result<ListResourcesResult, Error> {
|
||||
@@ -82,6 +71,7 @@ pub trait McpClientTrait: Send + Sync {
|
||||
|
||||
async fn read_resource(
|
||||
&self,
|
||||
_session_id: &str,
|
||||
_uri: &str,
|
||||
_cancel_token: CancellationToken,
|
||||
) -> Result<ReadResourceResult, Error> {
|
||||
@@ -90,6 +80,7 @@ pub trait McpClientTrait: Send + Sync {
|
||||
|
||||
async fn list_prompts(
|
||||
&self,
|
||||
_session_id: &str,
|
||||
_next_cursor: Option<String>,
|
||||
_cancel_token: CancellationToken,
|
||||
) -> Result<ListPromptsResult, Error> {
|
||||
@@ -98,6 +89,7 @@ pub trait McpClientTrait: Send + Sync {
|
||||
|
||||
async fn get_prompt(
|
||||
&self,
|
||||
_session_id: &str,
|
||||
_name: &str,
|
||||
_arguments: Value,
|
||||
_cancel_token: CancellationToken,
|
||||
@@ -117,6 +109,10 @@ pub trait McpClientTrait: Send + Sync {
|
||||
pub struct GooseClient {
|
||||
notification_handlers: Arc<Mutex<Vec<Sender<ServerNotification>>>>,
|
||||
provider: SharedProvider,
|
||||
// Single-slot because calls are serialized per MCP client; see send_request_with_session.
|
||||
current_session_id: Arc<Mutex<Option<String>>>,
|
||||
// Connection-scoped fallback for server-initiated sampling.
|
||||
client_session_id: OnceLock<String>,
|
||||
}
|
||||
|
||||
impl GooseClient {
|
||||
@@ -127,8 +123,52 @@ impl GooseClient {
|
||||
GooseClient {
|
||||
notification_handlers: handlers,
|
||||
provider,
|
||||
current_session_id: Arc::new(Mutex::new(None)),
|
||||
client_session_id: OnceLock::new(),
|
||||
}
|
||||
}
|
||||
|
||||
async fn set_current_session_id(&self, session_id: &str) {
|
||||
let mut slot = self.current_session_id.lock().await;
|
||||
*slot = Some(session_id.to_string());
|
||||
}
|
||||
|
||||
async fn clear_current_session_id(&self) {
|
||||
let mut slot = self.current_session_id.lock().await;
|
||||
*slot = None;
|
||||
}
|
||||
|
||||
async fn current_session_id(&self) -> Option<String> {
|
||||
let slot = self.current_session_id.lock().await;
|
||||
slot.clone()
|
||||
}
|
||||
|
||||
async fn resolve_session_id(&self, extensions: &Extensions) -> String {
|
||||
// Prefer explicit MCP metadata, then the active request scope.
|
||||
if let Some(session_id) = Self::session_id_from_extensions(extensions) {
|
||||
return session_id;
|
||||
}
|
||||
if let Some(session_id) = self.current_session_id().await {
|
||||
return session_id;
|
||||
}
|
||||
// Fallback for server-initiated sampling not tied to a request session.
|
||||
self.client_session_id()
|
||||
}
|
||||
|
||||
fn client_session_id(&self) -> String {
|
||||
self.client_session_id
|
||||
.get_or_init(|| Uuid::new_v4().to_string())
|
||||
.clone()
|
||||
}
|
||||
|
||||
fn session_id_from_extensions(extensions: &Extensions) -> Option<String> {
|
||||
let meta = extensions.get::<Meta>()?;
|
||||
meta.0
|
||||
.iter()
|
||||
.find(|(key, _)| key.eq_ignore_ascii_case(SESSION_ID_HEADER))
|
||||
.and_then(|(_, value)| value.as_str())
|
||||
.map(|value| value.to_string())
|
||||
}
|
||||
}
|
||||
|
||||
impl ClientHandler for GooseClient {
|
||||
@@ -175,7 +215,7 @@ impl ClientHandler for GooseClient {
|
||||
async fn create_message(
|
||||
&self,
|
||||
params: CreateMessageRequestParam,
|
||||
_context: RequestContext<RoleClient>,
|
||||
context: RequestContext<RoleClient>,
|
||||
) -> Result<CreateMessageResult, ErrorData> {
|
||||
let provider = self
|
||||
.provider
|
||||
@@ -189,6 +229,9 @@ impl ClientHandler for GooseClient {
|
||||
))?
|
||||
.clone();
|
||||
|
||||
// Prefer explicit MCP metadata, then the active request scope.
|
||||
let session_id = self.resolve_session_id(&context.extensions).await;
|
||||
|
||||
let provider_ready_messages: Vec<crate::conversation::message::Message> = params
|
||||
.messages
|
||||
.iter()
|
||||
@@ -211,7 +254,7 @@ impl ClientHandler for GooseClient {
|
||||
.unwrap_or("You are a general-purpose AI agent called goose");
|
||||
|
||||
let (response, usage) = provider
|
||||
.complete(system_prompt, &provider_ready_messages, &[])
|
||||
.complete(&session_id, system_prompt, &provider_ready_messages, &[])
|
||||
.await
|
||||
.map_err(|e| {
|
||||
ErrorData::new(
|
||||
@@ -336,19 +379,38 @@ impl McpClient {
|
||||
})
|
||||
}
|
||||
|
||||
async fn send_request(
|
||||
async fn send_request_with_session(
|
||||
&self,
|
||||
session_id: &str,
|
||||
request: ClientRequest,
|
||||
cancel_token: CancellationToken,
|
||||
) -> Result<ServerResult, Error> {
|
||||
let handle = self
|
||||
.client
|
||||
.lock()
|
||||
.await
|
||||
.send_cancellable_request(request, PeerRequestOptions::no_options())
|
||||
.await?;
|
||||
let request = inject_session_id_into_request(request, session_id);
|
||||
// ExtensionManager serializes calls per MCP connection, so one current_session_id slot
|
||||
// is sufficient for mapping callbacks to the active request session.
|
||||
let handle = {
|
||||
let client = self.client.lock().await;
|
||||
client.service().set_current_session_id(session_id).await;
|
||||
client
|
||||
.send_cancellable_request(request, PeerRequestOptions::no_options())
|
||||
.await
|
||||
};
|
||||
|
||||
await_response(handle, self.timeout, &cancel_token).await
|
||||
let handle = match handle {
|
||||
Ok(handle) => handle,
|
||||
Err(err) => {
|
||||
let client = self.client.lock().await;
|
||||
client.service().clear_current_session_id().await;
|
||||
return Err(err);
|
||||
}
|
||||
};
|
||||
|
||||
let result = await_response(handle, self.timeout, &cancel_token).await;
|
||||
|
||||
let client = self.client.lock().await;
|
||||
client.service().clear_current_session_id().await;
|
||||
|
||||
result
|
||||
}
|
||||
}
|
||||
|
||||
@@ -399,15 +461,17 @@ impl McpClientTrait for McpClient {
|
||||
|
||||
async fn list_resources(
|
||||
&self,
|
||||
session_id: &str,
|
||||
cursor: Option<String>,
|
||||
cancel_token: CancellationToken,
|
||||
) -> Result<ListResourcesResult, Error> {
|
||||
let res = self
|
||||
.send_request(
|
||||
.send_request_with_session(
|
||||
session_id,
|
||||
ClientRequest::ListResourcesRequest(ListResourcesRequest {
|
||||
params: Some(PaginatedRequestParam { cursor }),
|
||||
method: Default::default(),
|
||||
extensions: inject_current_session_id_into_extensions(Default::default()),
|
||||
extensions: Default::default(),
|
||||
}),
|
||||
cancel_token,
|
||||
)
|
||||
@@ -421,17 +485,19 @@ impl McpClientTrait for McpClient {
|
||||
|
||||
async fn read_resource(
|
||||
&self,
|
||||
session_id: &str,
|
||||
uri: &str,
|
||||
cancel_token: CancellationToken,
|
||||
) -> Result<ReadResourceResult, Error> {
|
||||
let res = self
|
||||
.send_request(
|
||||
.send_request_with_session(
|
||||
session_id,
|
||||
ClientRequest::ReadResourceRequest(ReadResourceRequest {
|
||||
params: ReadResourceRequestParam {
|
||||
uri: uri.to_string(),
|
||||
},
|
||||
method: Default::default(),
|
||||
extensions: inject_current_session_id_into_extensions(Default::default()),
|
||||
extensions: Default::default(),
|
||||
}),
|
||||
cancel_token,
|
||||
)
|
||||
@@ -445,15 +511,17 @@ impl McpClientTrait for McpClient {
|
||||
|
||||
async fn list_tools(
|
||||
&self,
|
||||
session_id: &str,
|
||||
cursor: Option<String>,
|
||||
cancel_token: CancellationToken,
|
||||
) -> Result<ListToolsResult, Error> {
|
||||
let res = self
|
||||
.send_request(
|
||||
.send_request_with_session(
|
||||
session_id,
|
||||
ClientRequest::ListToolsRequest(ListToolsRequest {
|
||||
params: Some(PaginatedRequestParam { cursor }),
|
||||
method: Default::default(),
|
||||
extensions: inject_current_session_id_into_extensions(Default::default()),
|
||||
extensions: Default::default(),
|
||||
}),
|
||||
cancel_token,
|
||||
)
|
||||
@@ -467,27 +535,26 @@ impl McpClientTrait for McpClient {
|
||||
|
||||
async fn call_tool(
|
||||
&self,
|
||||
session_id: &str,
|
||||
name: &str,
|
||||
arguments: Option<JsonObject>,
|
||||
meta: McpMeta,
|
||||
cancel_token: CancellationToken,
|
||||
) -> Result<CallToolResult, Error> {
|
||||
let res = self
|
||||
.send_request(
|
||||
ClientRequest::CallToolRequest(CallToolRequest {
|
||||
params: CallToolRequestParam {
|
||||
task: None,
|
||||
name: name.to_string().into(),
|
||||
arguments,
|
||||
},
|
||||
method: Default::default(),
|
||||
extensions: meta.inject_into_extensions(Default::default()),
|
||||
}),
|
||||
cancel_token,
|
||||
)
|
||||
.await?;
|
||||
let request = ClientRequest::CallToolRequest(CallToolRequest {
|
||||
params: CallToolRequestParam {
|
||||
task: None,
|
||||
name: name.to_string().into(),
|
||||
arguments,
|
||||
},
|
||||
method: Default::default(),
|
||||
extensions: Default::default(),
|
||||
});
|
||||
|
||||
match res {
|
||||
let result = self
|
||||
.send_request_with_session(session_id, request, cancel_token)
|
||||
.await;
|
||||
|
||||
match result? {
|
||||
ServerResult::CallToolResult(result) => Ok(result),
|
||||
_ => Err(ServiceError::UnexpectedResponse),
|
||||
}
|
||||
@@ -495,15 +562,17 @@ impl McpClientTrait for McpClient {
|
||||
|
||||
async fn list_prompts(
|
||||
&self,
|
||||
session_id: &str,
|
||||
cursor: Option<String>,
|
||||
cancel_token: CancellationToken,
|
||||
) -> Result<ListPromptsResult, Error> {
|
||||
let res = self
|
||||
.send_request(
|
||||
.send_request_with_session(
|
||||
session_id,
|
||||
ClientRequest::ListPromptsRequest(ListPromptsRequest {
|
||||
params: Some(PaginatedRequestParam { cursor }),
|
||||
method: Default::default(),
|
||||
extensions: inject_current_session_id_into_extensions(Default::default()),
|
||||
extensions: Default::default(),
|
||||
}),
|
||||
cancel_token,
|
||||
)
|
||||
@@ -517,6 +586,7 @@ impl McpClientTrait for McpClient {
|
||||
|
||||
async fn get_prompt(
|
||||
&self,
|
||||
session_id: &str,
|
||||
name: &str,
|
||||
arguments: Value,
|
||||
cancel_token: CancellationToken,
|
||||
@@ -526,14 +596,15 @@ impl McpClientTrait for McpClient {
|
||||
_ => None,
|
||||
};
|
||||
let res = self
|
||||
.send_request(
|
||||
.send_request_with_session(
|
||||
session_id,
|
||||
ClientRequest::GetPromptRequest(GetPromptRequest {
|
||||
params: GetPromptRequestParam {
|
||||
name: name.to_string(),
|
||||
arguments,
|
||||
},
|
||||
method: Default::default(),
|
||||
extensions: inject_current_session_id_into_extensions(Default::default()),
|
||||
extensions: Default::default(),
|
||||
}),
|
||||
cancel_token,
|
||||
)
|
||||
@@ -571,103 +642,241 @@ fn inject_session_id_into_extensions(mut extensions: Extensions, session_id: &st
|
||||
extensions
|
||||
}
|
||||
|
||||
/// Injects session ID from task-local context into Extensions._meta.
|
||||
fn inject_current_session_id_into_extensions(extensions: Extensions) -> Extensions {
|
||||
if let Some(session_id) = crate::session_context::current_session_id() {
|
||||
inject_session_id_into_extensions(extensions, &session_id)
|
||||
} else {
|
||||
extensions
|
||||
fn inject_session_id_into_request(request: ClientRequest, session_id: &str) -> ClientRequest {
|
||||
match request {
|
||||
ClientRequest::ListResourcesRequest(mut req) => {
|
||||
req.extensions = inject_session_id_into_extensions(req.extensions, session_id);
|
||||
ClientRequest::ListResourcesRequest(req)
|
||||
}
|
||||
ClientRequest::ReadResourceRequest(mut req) => {
|
||||
req.extensions = inject_session_id_into_extensions(req.extensions, session_id);
|
||||
ClientRequest::ReadResourceRequest(req)
|
||||
}
|
||||
ClientRequest::ListToolsRequest(mut req) => {
|
||||
req.extensions = inject_session_id_into_extensions(req.extensions, session_id);
|
||||
ClientRequest::ListToolsRequest(req)
|
||||
}
|
||||
ClientRequest::CallToolRequest(mut req) => {
|
||||
req.extensions = inject_session_id_into_extensions(req.extensions, session_id);
|
||||
ClientRequest::CallToolRequest(req)
|
||||
}
|
||||
ClientRequest::ListPromptsRequest(mut req) => {
|
||||
req.extensions = inject_session_id_into_extensions(req.extensions, session_id);
|
||||
ClientRequest::ListPromptsRequest(req)
|
||||
}
|
||||
ClientRequest::GetPromptRequest(mut req) => {
|
||||
req.extensions = inject_session_id_into_extensions(req.extensions, session_id);
|
||||
ClientRequest::GetPromptRequest(req)
|
||||
}
|
||||
other => other,
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use rmcp::model::Meta;
|
||||
use test_case::test_case;
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_session_id_in_mcp_meta() {
|
||||
fn new_client() -> GooseClient {
|
||||
GooseClient::new(Arc::new(Mutex::new(Vec::new())), Arc::new(Mutex::new(None)))
|
||||
}
|
||||
|
||||
fn request_extensions(request: &ClientRequest) -> Option<&Extensions> {
|
||||
match request {
|
||||
ClientRequest::ListResourcesRequest(req) => Some(&req.extensions),
|
||||
ClientRequest::ReadResourceRequest(req) => Some(&req.extensions),
|
||||
ClientRequest::ListToolsRequest(req) => Some(&req.extensions),
|
||||
ClientRequest::CallToolRequest(req) => Some(&req.extensions),
|
||||
ClientRequest::ListPromptsRequest(req) => Some(&req.extensions),
|
||||
ClientRequest::GetPromptRequest(req) => Some(&req.extensions),
|
||||
_ => None,
|
||||
}
|
||||
}
|
||||
|
||||
fn list_resources_request(extensions: Extensions) -> ClientRequest {
|
||||
ClientRequest::ListResourcesRequest(ListResourcesRequest {
|
||||
params: Some(PaginatedRequestParam { cursor: None }),
|
||||
method: Default::default(),
|
||||
extensions,
|
||||
})
|
||||
}
|
||||
|
||||
fn read_resource_request(extensions: Extensions) -> ClientRequest {
|
||||
ClientRequest::ReadResourceRequest(ReadResourceRequest {
|
||||
params: ReadResourceRequestParam {
|
||||
uri: "test://resource".to_string(),
|
||||
},
|
||||
method: Default::default(),
|
||||
extensions,
|
||||
})
|
||||
}
|
||||
|
||||
fn list_tools_request(extensions: Extensions) -> ClientRequest {
|
||||
ClientRequest::ListToolsRequest(ListToolsRequest {
|
||||
params: Some(PaginatedRequestParam { cursor: None }),
|
||||
method: Default::default(),
|
||||
extensions,
|
||||
})
|
||||
}
|
||||
|
||||
fn call_tool_request(extensions: Extensions) -> ClientRequest {
|
||||
ClientRequest::CallToolRequest(CallToolRequest {
|
||||
params: CallToolRequestParam {
|
||||
task: None,
|
||||
name: "tool".to_string().into(),
|
||||
arguments: None,
|
||||
},
|
||||
method: Default::default(),
|
||||
extensions,
|
||||
})
|
||||
}
|
||||
|
||||
fn list_prompts_request(extensions: Extensions) -> ClientRequest {
|
||||
ClientRequest::ListPromptsRequest(ListPromptsRequest {
|
||||
params: Some(PaginatedRequestParam { cursor: None }),
|
||||
method: Default::default(),
|
||||
extensions,
|
||||
})
|
||||
}
|
||||
|
||||
fn get_prompt_request(extensions: Extensions) -> ClientRequest {
|
||||
ClientRequest::GetPromptRequest(GetPromptRequest {
|
||||
params: GetPromptRequestParam {
|
||||
name: "prompt".to_string(),
|
||||
arguments: None,
|
||||
},
|
||||
method: Default::default(),
|
||||
extensions,
|
||||
})
|
||||
}
|
||||
|
||||
#[test_case(
|
||||
Some("ext-session"),
|
||||
Some("current-session"),
|
||||
"ext-session";
|
||||
"extensions win"
|
||||
)]
|
||||
#[test_case(
|
||||
None,
|
||||
Some("current-session"),
|
||||
"current-session";
|
||||
"current when no extensions"
|
||||
)]
|
||||
#[test_case(
|
||||
None,
|
||||
None,
|
||||
"client-session";
|
||||
"client fallback when no session"
|
||||
)]
|
||||
fn test_resolve_session_id(
|
||||
ext_session: Option<&str>,
|
||||
current_session: Option<&str>,
|
||||
expected: &str,
|
||||
) {
|
||||
let runtime = tokio::runtime::Runtime::new().unwrap();
|
||||
runtime.block_on(async {
|
||||
let client = new_client();
|
||||
// Make the fallback deterministic so the expected value can live in the test_case row.
|
||||
client
|
||||
.client_session_id
|
||||
.get_or_init(|| "client-session".to_string());
|
||||
if let Some(session_id) = current_session {
|
||||
let mut slot = client.current_session_id.lock().await;
|
||||
*slot = Some(session_id.to_string());
|
||||
}
|
||||
|
||||
let mut extensions = Extensions::new();
|
||||
if let Some(session_id) = ext_session {
|
||||
extensions = inject_session_id_into_extensions(extensions, session_id);
|
||||
}
|
||||
|
||||
let resolved = client.resolve_session_id(&extensions).await;
|
||||
|
||||
assert_eq!(resolved, expected);
|
||||
});
|
||||
}
|
||||
|
||||
#[test_case(list_resources_request; "list_resources")]
|
||||
#[test_case(read_resource_request; "read_resource")]
|
||||
#[test_case(list_tools_request; "list_tools")]
|
||||
#[test_case(call_tool_request; "call_tool")]
|
||||
#[test_case(list_prompts_request; "list_prompts")]
|
||||
#[test_case(get_prompt_request; "get_prompt")]
|
||||
fn test_request_injects_session(request_builder: fn(Extensions) -> ClientRequest) {
|
||||
use serde_json::json;
|
||||
|
||||
let session_id = "test-session-id";
|
||||
let mut extensions = Extensions::new();
|
||||
extensions.insert(
|
||||
serde_json::from_value::<Meta>(json!({
|
||||
"Goose-Session-Id": "old-session-id",
|
||||
"other-key": "preserve-me"
|
||||
}))
|
||||
.unwrap(),
|
||||
);
|
||||
|
||||
let request = request_builder(extensions);
|
||||
let request = inject_session_id_into_request(request, session_id);
|
||||
let extensions = request_extensions(&request).expect("request should have extensions");
|
||||
let meta = extensions
|
||||
.get::<Meta>()
|
||||
.expect("extensions should contain meta");
|
||||
|
||||
assert_eq!(
|
||||
meta.0.get(SESSION_ID_HEADER),
|
||||
Some(&Value::String(session_id.to_string()))
|
||||
);
|
||||
assert_eq!(
|
||||
meta.0.get("other-key"),
|
||||
Some(&Value::String("preserve-me".to_string()))
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_session_id_in_mcp_meta() {
|
||||
use serde_json::json;
|
||||
|
||||
let session_id = "test-session-789";
|
||||
crate::session_context::with_session_id(Some(session_id.to_string()), async {
|
||||
let extensions = inject_current_session_id_into_extensions(Default::default());
|
||||
let meta = extensions.get::<Meta>().unwrap();
|
||||
let extensions = inject_session_id_into_extensions(Default::default(), session_id);
|
||||
let mcp_meta = extensions.get::<Meta>().unwrap();
|
||||
|
||||
assert_eq!(
|
||||
&meta.0,
|
||||
json!({
|
||||
SESSION_ID_HEADER: session_id
|
||||
})
|
||||
.as_object()
|
||||
.unwrap()
|
||||
);
|
||||
})
|
||||
.await;
|
||||
assert_eq!(
|
||||
&mcp_meta.0,
|
||||
json!({
|
||||
SESSION_ID_HEADER: session_id
|
||||
})
|
||||
.as_object()
|
||||
.unwrap()
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_no_session_id_in_mcp_when_absent() {
|
||||
let extensions = inject_current_session_id_into_extensions(Default::default());
|
||||
let meta = extensions.get::<Meta>();
|
||||
|
||||
assert!(meta.is_none());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_all_mcp_operations_include_session() {
|
||||
use serde_json::json;
|
||||
|
||||
let session_id = "consistent-session-id";
|
||||
crate::session_context::with_session_id(Some(session_id.to_string()), async {
|
||||
let ext1 = inject_current_session_id_into_extensions(Default::default());
|
||||
let ext2 = inject_current_session_id_into_extensions(Default::default());
|
||||
let ext3 = inject_current_session_id_into_extensions(Default::default());
|
||||
|
||||
for ext in [&ext1, &ext2, &ext3] {
|
||||
assert_eq!(
|
||||
&ext.get::<Meta>().unwrap().0,
|
||||
json!({
|
||||
SESSION_ID_HEADER: session_id
|
||||
})
|
||||
.as_object()
|
||||
.unwrap()
|
||||
);
|
||||
}
|
||||
})
|
||||
.await;
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_session_id_case_insensitive_replacement() {
|
||||
use rmcp::model::{Extensions, Meta};
|
||||
#[test]
|
||||
fn test_session_id_case_insensitive_replacement() {
|
||||
use rmcp::model::Extensions;
|
||||
use serde_json::{from_value, json};
|
||||
|
||||
let session_id = "new-session-id";
|
||||
crate::session_context::with_session_id(Some(session_id.to_string()), async {
|
||||
let mut extensions = Extensions::new();
|
||||
extensions.insert(
|
||||
from_value::<Meta>(json!({
|
||||
"GOOSE-SESSION-ID": "old-session-1",
|
||||
"Goose-Session-Id": "old-session-2",
|
||||
"other-key": "preserve-me"
|
||||
}))
|
||||
.unwrap(),
|
||||
);
|
||||
let mut extensions = Extensions::new();
|
||||
extensions.insert(
|
||||
from_value::<Meta>(json!({
|
||||
"GOOSE-SESSION-ID": "old-session-1",
|
||||
"Goose-Session-Id": "old-session-2",
|
||||
"other-key": "preserve-me"
|
||||
}))
|
||||
.unwrap(),
|
||||
);
|
||||
|
||||
let extensions = inject_current_session_id_into_extensions(extensions);
|
||||
let meta = extensions.get::<Meta>().unwrap();
|
||||
let extensions = inject_session_id_into_extensions(extensions, session_id);
|
||||
let mcp_meta = extensions.get::<Meta>().unwrap();
|
||||
|
||||
assert_eq!(
|
||||
&meta.0,
|
||||
json!({
|
||||
SESSION_ID_HEADER: session_id,
|
||||
"other-key": "preserve-me"
|
||||
})
|
||||
.as_object()
|
||||
.unwrap()
|
||||
);
|
||||
})
|
||||
.await;
|
||||
assert_eq!(
|
||||
&mcp_meta.0,
|
||||
json!({
|
||||
SESSION_ID_HEADER: session_id,
|
||||
"other-key": "preserve-me"
|
||||
})
|
||||
.as_object()
|
||||
.unwrap()
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user