fix(goose): only send agent-session-id when a session exists (#6657)
Signed-off-by: Adrian Cole <adrian@tetrate.io>
This commit is contained in:
@@ -27,17 +27,12 @@ use rmcp::{
|
||||
ClientHandler, ErrorData, Peer, RoleClient, ServiceError, ServiceExt,
|
||||
};
|
||||
use serde_json::Value;
|
||||
use std::{
|
||||
sync::{Arc, OnceLock},
|
||||
time::Duration,
|
||||
};
|
||||
use std::{sync::Arc, 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;
|
||||
@@ -112,8 +107,6 @@ pub struct GooseClient {
|
||||
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 {
|
||||
@@ -125,7 +118,6 @@ impl GooseClient {
|
||||
notification_handlers: handlers,
|
||||
provider,
|
||||
current_session_id: Arc::new(Mutex::new(None)),
|
||||
client_session_id: OnceLock::new(),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -144,22 +136,10 @@ impl GooseClient {
|
||||
slot.clone()
|
||||
}
|
||||
|
||||
async fn resolve_session_id(&self, extensions: &Extensions) -> String {
|
||||
async fn resolve_session_id(&self, extensions: &Extensions) -> Option<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()
|
||||
let current_session_id = self.current_session_id().await;
|
||||
Self::session_id_from_extensions(extensions).or(current_session_id)
|
||||
}
|
||||
|
||||
fn session_id_from_extensions(extensions: &Extensions) -> Option<String> {
|
||||
@@ -255,7 +235,13 @@ impl ClientHandler for GooseClient {
|
||||
.unwrap_or("You are a general-purpose AI agent called goose");
|
||||
|
||||
let (response, usage) = provider
|
||||
.complete(&session_id, system_prompt, &provider_ready_messages, &[])
|
||||
.complete_with_model(
|
||||
session_id.as_deref(),
|
||||
&provider.get_model_config(),
|
||||
system_prompt,
|
||||
&provider_ready_messages,
|
||||
&[],
|
||||
)
|
||||
.await
|
||||
.map_err(|e| {
|
||||
ErrorData::new(
|
||||
@@ -387,7 +373,7 @@ impl McpClient {
|
||||
request: ClientRequest,
|
||||
cancel_token: CancellationToken,
|
||||
) -> Result<ServerResult, Error> {
|
||||
let request = inject_session_id_into_request(request, session_id);
|
||||
let request = inject_session_id_into_request(request, Some(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 = {
|
||||
@@ -629,7 +615,12 @@ impl McpClientTrait for McpClient {
|
||||
}
|
||||
|
||||
/// Injects the given session_id into Extensions._meta.
|
||||
fn inject_session_id_into_extensions(mut extensions: Extensions, session_id: &str) -> Extensions {
|
||||
/// None (or empty) removes any existing session id.
|
||||
fn inject_session_id_into_extensions(
|
||||
mut extensions: Extensions,
|
||||
session_id: Option<&str>,
|
||||
) -> Extensions {
|
||||
let session_id = session_id.filter(|id| !id.is_empty());
|
||||
let mut meta_map = extensions
|
||||
.get::<Meta>()
|
||||
.map(|meta| meta.0.clone())
|
||||
@@ -638,16 +629,21 @@ fn inject_session_id_into_extensions(mut extensions: Extensions, session_id: &st
|
||||
// JsonObject is case-sensitive, so we use retain for case-insensitive removal
|
||||
meta_map.retain(|k, _| !k.eq_ignore_ascii_case(SESSION_ID_HEADER));
|
||||
|
||||
meta_map.insert(
|
||||
SESSION_ID_HEADER.to_string(),
|
||||
Value::String(session_id.to_string()),
|
||||
);
|
||||
if let Some(session_id) = session_id {
|
||||
meta_map.insert(
|
||||
SESSION_ID_HEADER.to_string(),
|
||||
Value::String(session_id.to_string()),
|
||||
);
|
||||
}
|
||||
|
||||
extensions.insert(Meta(meta_map));
|
||||
extensions
|
||||
}
|
||||
|
||||
fn inject_session_id_into_request(request: ClientRequest, session_id: &str) -> ClientRequest {
|
||||
fn inject_session_id_into_request(
|
||||
request: ClientRequest,
|
||||
session_id: Option<&str>,
|
||||
) -> ClientRequest {
|
||||
match request {
|
||||
ClientRequest::ListResourcesRequest(mut req) => {
|
||||
req.extensions = inject_session_id_into_extensions(req.extensions, session_id);
|
||||
@@ -680,6 +676,7 @@ fn inject_session_id_into_request(request: ClientRequest, session_id: &str) -> C
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use serde_json::json;
|
||||
use test_case::test_case;
|
||||
|
||||
fn new_client() -> GooseClient {
|
||||
@@ -770,45 +767,39 @@ mod tests {
|
||||
#[test_case(
|
||||
Some("ext-session"),
|
||||
Some("current-session"),
|
||||
"ext-session";
|
||||
Some("ext-session");
|
||||
"extensions win"
|
||||
)]
|
||||
#[test_case(
|
||||
None,
|
||||
Some("current-session"),
|
||||
"current-session";
|
||||
Some("current-session");
|
||||
"current when no extensions"
|
||||
)]
|
||||
#[test_case(
|
||||
None,
|
||||
None,
|
||||
"client-session";
|
||||
"client fallback when no session"
|
||||
None;
|
||||
"no session when no extensions or current"
|
||||
)]
|
||||
fn test_resolve_session_id(
|
||||
ext_session: Option<&str>,
|
||||
current_session: Option<&str>,
|
||||
expected: &str,
|
||||
expected: Option<&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 extensions = inject_session_id_into_extensions(Extensions::new(), ext_session);
|
||||
|
||||
let resolved = client.resolve_session_id(&extensions).await;
|
||||
|
||||
let expected = expected.map(str::to_string);
|
||||
assert_eq!(resolved, expected);
|
||||
});
|
||||
}
|
||||
@@ -833,7 +824,7 @@ mod tests {
|
||||
);
|
||||
|
||||
let request = request_builder(extensions);
|
||||
let request = inject_session_id_into_request(request, session_id);
|
||||
let request = inject_session_id_into_request(request, Some(session_id));
|
||||
let extensions = request_extensions(&request).expect("request should have extensions");
|
||||
let meta = extensions
|
||||
.get::<Meta>()
|
||||
@@ -854,7 +845,7 @@ mod tests {
|
||||
use serde_json::json;
|
||||
|
||||
let session_id = "test-session-789";
|
||||
let extensions = inject_session_id_into_extensions(Default::default(), session_id);
|
||||
let extensions = inject_session_id_into_extensions(Default::default(), Some(session_id));
|
||||
let mcp_meta = extensions.get::<Meta>().unwrap();
|
||||
|
||||
assert_eq!(
|
||||
@@ -867,12 +858,35 @@ mod tests {
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_session_id_case_insensitive_replacement() {
|
||||
#[test_case(
|
||||
Some("new-session-id"),
|
||||
json!({
|
||||
SESSION_ID_HEADER: "new-session-id",
|
||||
"other-key": "preserve-me"
|
||||
});
|
||||
"replace"
|
||||
)]
|
||||
#[test_case(
|
||||
None,
|
||||
json!({
|
||||
"other-key": "preserve-me"
|
||||
});
|
||||
"remove"
|
||||
)]
|
||||
#[test_case(
|
||||
Some(""),
|
||||
json!({
|
||||
"other-key": "preserve-me"
|
||||
});
|
||||
"empty removes"
|
||||
)]
|
||||
fn test_session_id_case_insensitive_replacement(
|
||||
session_id: Option<&str>,
|
||||
expected_meta: serde_json::Value,
|
||||
) {
|
||||
use rmcp::model::Extensions;
|
||||
use serde_json::{from_value, json};
|
||||
|
||||
let session_id = "new-session-id";
|
||||
let mut extensions = Extensions::new();
|
||||
extensions.insert(
|
||||
from_value::<Meta>(json!({
|
||||
@@ -886,14 +900,6 @@ mod tests {
|
||||
let extensions = inject_session_id_into_extensions(extensions, session_id);
|
||||
let mcp_meta = extensions.get::<Meta>().unwrap();
|
||||
|
||||
assert_eq!(
|
||||
&mcp_meta.0,
|
||||
json!({
|
||||
SESSION_ID_HEADER: session_id,
|
||||
"other-key": "preserve-me"
|
||||
})
|
||||
.as_object()
|
||||
.unwrap()
|
||||
);
|
||||
assert_eq!(&mcp_meta.0, expected_meta.as_object().unwrap());
|
||||
}
|
||||
}
|
||||
|
||||
@@ -435,7 +435,7 @@ mod tests {
|
||||
|
||||
async fn complete_with_model(
|
||||
&self,
|
||||
_session_id: &str,
|
||||
_session_id: Option<&str>,
|
||||
_model_config: &ModelConfig,
|
||||
_system: &str,
|
||||
_messages: &[Message],
|
||||
|
||||
Reference in New Issue
Block a user