fix(goose): propagate session_id across providers and MCP (#6584)

Signed-off-by: Adrian Cole <adrian@tetrate.io>
This commit is contained in:
Adrian Cole
2026-01-22 09:28:56 +09:00
committed by GitHub
parent f3bae7ea7a
commit 67de49abbb
61 changed files with 1457 additions and 616 deletions
+70 -42
View File
@@ -196,6 +196,7 @@ pub struct ApiRequestBuilder<'a> {
client: &'a ApiClient,
path: &'a str,
headers: HeaderMap,
session_id: &'a str,
}
impl ApiClient {
@@ -272,28 +273,39 @@ impl ApiClient {
Ok(self)
}
pub fn request<'a>(&'a self, path: &'a str) -> ApiRequestBuilder<'a> {
pub fn request<'a>(&'a self, session_id: &'a str, path: &'a str) -> ApiRequestBuilder<'a> {
ApiRequestBuilder {
client: self,
session_id,
path,
headers: HeaderMap::new(),
}
}
pub async fn api_post(&self, path: &str, payload: &Value) -> Result<ApiResponse> {
self.request(path).api_post(payload).await
pub async fn api_post(
&self,
session_id: &str,
path: &str,
payload: &Value,
) -> Result<ApiResponse> {
self.request(session_id, path).api_post(payload).await
}
pub async fn response_post(&self, path: &str, payload: &Value) -> Result<Response> {
self.request(path).response_post(payload).await
pub async fn response_post(
&self,
session_id: &str,
path: &str,
payload: &Value,
) -> Result<Response> {
self.request(session_id, path).response_post(payload).await
}
pub async fn api_get(&self, path: &str) -> Result<ApiResponse> {
self.request(path).api_get().await
pub async fn api_get(&self, session_id: &str, path: &str) -> Result<ApiResponse> {
self.request(session_id, path).api_get().await
}
pub async fn response_get(&self, path: &str) -> Result<Response> {
self.request(path).response_get().await
pub async fn response_get(&self, session_id: &str, path: &str) -> Result<Response> {
self.request(session_id, path).response_get().await
}
fn build_url(&self, path: &str) -> Result<url::Url> {
@@ -370,9 +382,7 @@ impl<'a> ApiRequestBuilder<'a> {
let mut request = request_builder(url, &self.client.client);
request = request.headers(self.headers.clone());
if let Some(session_id) = crate::session_context::current_session_id() {
request = request.header(SESSION_ID_HEADER, session_id);
}
request = request.header(SESSION_ID_HEADER, self.session_id);
request = match &self.client.auth {
AuthMethod::BearerToken(token) => {
@@ -416,35 +426,7 @@ mod tests {
)
.unwrap();
// Execute request within session context
crate::session_context::with_session_id(Some("test-session-456".to_string()), async {
let builder = client.request("/test");
let request = builder
.send_request(|url, client| client.get(url))
.await
.unwrap();
let headers = request.build().unwrap().headers().clone();
assert!(headers.contains_key(SESSION_ID_HEADER));
assert_eq!(
headers.get(SESSION_ID_HEADER).unwrap().to_str().unwrap(),
"test-session-456"
);
})
.await;
}
#[tokio::test]
async fn test_no_session_id_header_when_absent() {
let client = ApiClient::new(
"http://localhost:8080".to_string(),
AuthMethod::BearerToken("test-token".to_string()),
)
.unwrap();
// Build a request without session context
let builder = client.request("/test");
let builder = client.request("test-session_id-456", "/test");
let request = builder
.send_request(|url, client| client.get(url))
.await
@@ -452,6 +434,52 @@ mod tests {
let headers = request.build().unwrap().headers().clone();
assert!(!headers.contains_key(SESSION_ID_HEADER));
assert!(headers.contains_key(SESSION_ID_HEADER));
assert_eq!(
headers.get(SESSION_ID_HEADER).unwrap().to_str().unwrap(),
"test-session_id-456"
);
}
#[tokio::test]
async fn test_session_id_header_with_different_id() {
let client = ApiClient::new(
"http://localhost:8080".to_string(),
AuthMethod::BearerToken("test-token".to_string()),
)
.unwrap();
let builder = client.request("another-session_id-789", "/test");
let request = builder
.send_request(|url, client| client.get(url))
.await
.unwrap();
let headers = request.build().unwrap().headers().clone();
assert!(headers.contains_key(SESSION_ID_HEADER));
assert_eq!(
headers.get(SESSION_ID_HEADER).unwrap().to_str().unwrap(),
"another-session_id-789"
);
}
#[tokio::test]
async fn test_session_id_header_always_present() {
let client = ApiClient::new(
"http://localhost:8080".to_string(),
AuthMethod::BearerToken("test-token".to_string()),
)
.unwrap();
let builder = client.request("required-session_id", "/test");
let request = builder
.send_request(|url, client| client.get(url))
.await
.unwrap();
let headers = request.build().unwrap().headers().clone();
assert!(headers.contains_key(SESSION_ID_HEADER));
}
}