diff --git a/crates/goose-cli/src/commands/configure.rs b/crates/goose-cli/src/commands/configure.rs index d94b1d6e5e23..bb51673d1af3 100644 --- a/crates/goose-cli/src/commands/configure.rs +++ b/crates/goose-cli/src/commands/configure.rs @@ -1813,7 +1813,6 @@ pub async fn handle_openrouter_auth() -> anyhow::Result<()> { let test_result = provider .complete( &model_config, - "", "You are goose, an AI assistant.", &[Message::user().with_text("Say 'Configuration test successful!'")], &[], diff --git a/crates/goose-cli/src/commands/info.rs b/crates/goose-cli/src/commands/info.rs index d21771d29cb5..f471018cb346 100644 --- a/crates/goose-cli/src/commands/info.rs +++ b/crates/goose-cli/src/commands/info.rs @@ -88,10 +88,12 @@ async fn check_provider( let test_msg = Message::user().with_text("Say 'ok'"); let start = std::time::Instant::now(); - provider_client - .complete(&model_config, "check", "", &[test_msg], &[]) - .await - .map_err(ProviderCheckError::ProviderRequest)?; + goose::session_context::with_session_id( + Some("check".to_string()), + provider_client.complete(&model_config, "", &[test_msg], &[]), + ) + .await + .map_err(ProviderCheckError::ProviderRequest)?; Ok(ProviderCheckSuccess { provider, diff --git a/crates/goose-cli/src/session/mod.rs b/crates/goose-cli/src/session/mod.rs index 45e4acf06093..73dfe727237f 100644 --- a/crates/goose-cli/src/session/mod.rs +++ b/crates/goose-cli/src/session/mod.rs @@ -226,15 +226,16 @@ pub async fn classify_planner_response( ); let message = Message::user().with_text(&prompt); - let (result, _usage) = provider - .complete( + let (result, _usage) = goose::session_context::with_session_id( + Some(session_id.to_string()), + provider.complete( &model_config, - session_id, "Reply only with the classification label: \"plan\" or \"clarifying questions\"", &[message], &[], - ) - .await?; + ), + ) + .await?; let predicted = result.as_concat_text(); if predicted.to_lowercase().contains("plan") { @@ -1060,15 +1061,11 @@ impl CliSession { ) -> Result<(), anyhow::Error> { let plan_prompt = self.agent.get_plan_prompt(&self.session_id).await?; output::show_thinking(); - let (plan_response, _usage) = reasoner - .complete( - &model_config, - &self.session_id, - &plan_prompt, - plan_messages.messages(), - &[], - ) - .await?; + let (plan_response, _usage) = goose::session_context::with_session_id( + Some(self.session_id.clone()), + reasoner.complete(&model_config, &plan_prompt, plan_messages.messages(), &[]), + ) + .await?; output::render_message(&plan_response, self.debug); output::hide_thinking(); let planner_response_type = classify_planner_response( diff --git a/crates/goose-providers/examples/streaming.rs b/crates/goose-providers/examples/streaming.rs index efb482be3a5f..66ea1da7373e 100644 --- a/crates/goose-providers/examples/streaming.rs +++ b/crates/goose-providers/examples/streaming.rs @@ -10,6 +10,19 @@ use goose_providers::{ openai::OpenAiProvider, }; +async fn stream(provider: impl Provider, model: ModelConfig) -> Result<()> { + let system = "You are a knowledgable geography expert"; + let messages = [Message::user().with_text("what is the capital of France?")]; + + let mut stream = provider.stream(&model, system, &messages, &[]).await?; + + while let Some((Some(msg), _)) = stream.next().await.transpose()? { + print!("{}", msg.as_concat_text()); + } + println!(); + Ok(()) +} + #[tokio::main] async fn main() -> Result<()> { let key = env::var("OPENAI_API_KEY").map_err(|_| anyhow::anyhow!("need an OpenAI key"))?; @@ -18,26 +31,9 @@ async fn main() -> Result<()> { AuthMethod::BearerToken(key), Some(Default::default()), )?; - let provider = OpenAiProvider::new(api_client); - - let system = "You are a knowledgable geography expert"; - let messages = [Message::user().with_text("what is the capital of France?")]; + let provider = OpenAiProvider::new(api_client); let model = ModelConfig::new("gpt-5.4-mini"); - let mut stream = provider - .stream( - &model, - "", // session-id - system, - &messages, - &[], - ) - .await?; - while let Some((Some(msg), _)) = stream.next().await.transpose()? { - print!("{}", msg.as_concat_text()); - } - println!(); - - Ok(()) + stream(provider, model).await } diff --git a/crates/goose-providers/src/anthropic.rs b/crates/goose-providers/src/anthropic.rs index 364fa1e46b86..a0c901721c76 100644 --- a/crates/goose-providers/src/anthropic.rs +++ b/crates/goose-providers/src/anthropic.rs @@ -135,7 +135,7 @@ impl AnthropicProviderBuilder { impl AnthropicProvider { async fn fetch_models_from_api(&self) -> Result, ProviderError> { - let response = self.api_client.request(None, "v1/models").api_get().await?; + let response = self.api_client.request("v1/models").api_get().await?; if response.status == StatusCode::NOT_FOUND { let msg = response @@ -241,7 +241,6 @@ impl Provider for AnthropicProvider { async fn stream( &self, model_config: &ModelConfig, - session_id: &str, system: &str, messages: &[Message], tools: &[Tool], @@ -263,7 +262,7 @@ impl Provider for AnthropicProvider { let response = self .with_retry(|| async { - let request = self.api_client.request(Some(session_id), "v1/messages"); + let request = self.api_client.request("v1/messages"); let resp = request.response_post(&payload).await?; handle_status(resp).await }) diff --git a/crates/goose-providers/src/api_client.rs b/crates/goose-providers/src/api_client.rs index 2493e6ada7c4..34a6d7f02140 100644 --- a/crates/goose-providers/src/api_client.rs +++ b/crates/goose-providers/src/api_client.rs @@ -11,10 +11,13 @@ use std::fmt; #[cfg(any(feature = "rustls-tls", feature = "native-tls"))] use std::fs::read_to_string; use std::path::PathBuf; +use std::sync::Arc; use std::time::Duration; const DEFAULT_PROVIDER_TIMEOUT_SECS: u64 = 600; -const SESSION_ID_HEADER: &str = "agent-session-id"; + +pub type RequestBuilderDecorator = + Arc Result + Send + Sync>; pub struct ApiClient { client: Client, @@ -24,6 +27,7 @@ pub struct ApiClient { default_query: Vec<(String, String)>, timeout: Duration, tls_config: Option, + request_builder: Option, } pub enum AuthMethod { @@ -225,7 +229,6 @@ pub struct ApiRequestBuilder<'a> { client: &'a ApiClient, path: &'a str, headers: HeaderMap, - session_id: Option<&'a str>, } impl ApiClient { @@ -264,6 +267,7 @@ impl ApiClient { default_query: Vec::new(), timeout, tls_config, + request_builder: None, }) } @@ -339,44 +343,33 @@ impl ApiClient { Ok(self) } - /// - `session_id`: Use `None` only for configuration or pre-session tasks. - pub fn request<'a>( - &'a self, - session_id: Option<&'a str>, - path: &'a str, - ) -> ApiRequestBuilder<'a> { + pub fn with_request_builder(mut self, request_builder: RequestBuilderDecorator) -> Self { + self.request_builder = Some(request_builder); + self + } + + pub fn request<'a>(&'a self, path: &'a str) -> ApiRequestBuilder<'a> { ApiRequestBuilder { client: self, - session_id: session_id.filter(|id| !id.is_empty()), path, headers: HeaderMap::new(), } } - pub async fn api_post( - &self, - session_id: Option<&str>, - path: &str, - payload: &Value, - ) -> Result { - self.request(session_id, path).api_post(payload).await + pub async fn api_post(&self, path: &str, payload: &Value) -> Result { + self.request(path).api_post(payload).await } - pub async fn response_post( - &self, - session_id: Option<&str>, - path: &str, - payload: &Value, - ) -> Result { - self.request(session_id, path).response_post(payload).await + pub async fn response_post(&self, path: &str, payload: &Value) -> Result { + self.request(path).response_post(payload).await } - pub async fn api_get(&self, session_id: Option<&str>, path: &str) -> Result { - self.request(session_id, path).api_get().await + pub async fn api_get(&self, path: &str) -> Result { + self.request(path).api_get().await } - pub async fn response_get(&self, session_id: Option<&str>, path: &str) -> Result { - self.request(session_id, path).response_get().await + pub async fn response_get(&self, path: &str) -> Result { + self.request(path).response_get().await } fn build_url(&self, path: &str) -> Result { @@ -445,17 +438,14 @@ impl<'a> ApiRequestBuilder<'a> { F: FnOnce(url::Url, &Client) -> reqwest::RequestBuilder, { let url = self.client.build_url(self.path)?; - let mut headers = self.headers.clone(); - headers.remove(SESSION_ID_HEADER); - if let Some(session_id) = self.session_id { - let header_name = HeaderName::from_static(SESSION_ID_HEADER); - let header_value = HeaderValue::from_str(session_id)?; - headers.insert(header_name, header_value); - } - + let headers = self.headers.clone(); let mut request = request_builder(url, &self.client.client); request = request.headers(headers); + if let Some(decorator) = &self.client.request_builder { + request = decorator(request)?; + } + request = match &self.client.auth { AuthMethod::NoAuth => request, AuthMethod::BearerToken(token) => { @@ -607,17 +597,9 @@ ShGoCNbfNS+COlPMRAujyDlATZcLs9p4tA== #[cfg(test)] mod tests { use super::*; - use test_case::test_case; - - #[test_case(Some("test-session_id-456"), None, Some("test-session_id-456"); "header set")] - #[test_case(Some("new-session"), Some(("Agent-Session-Id", "old-session")), Some("new-session"); "replaces existing")] - #[test_case(None, Some(("Agent-Session-Id", "old-session")), None; "removes existing on none")] - #[test_case(Some(""), Some(("agent-session-id", "old-session")), None; "removes existing on empty")] - fn test_session_id_header( - session_id: Option<&str>, - existing_header: Option<(&str, &str)>, - expected: Option<&str>, - ) { + + #[test] + fn test_request_builder_decorator() { let runtime = tokio::runtime::Runtime::new().unwrap(); runtime.block_on(async { let client = ApiClient::new_with_tls( @@ -625,23 +607,22 @@ mod tests { AuthMethod::BearerToken("test-token".to_string()), None, ) - .unwrap(); + .unwrap() + .with_request_builder(Arc::new(|request| { + Ok(request.header("test-my-session-id", "test-session_id-456")) + })); - let mut builder = client.request(session_id, "/test"); - if let Some((key, value)) = existing_header { - builder = builder.header(key, value).unwrap(); - } - let request = builder + let request = client + .request("/test") .send_request(|url, client| client.get(url)) .await .unwrap(); let headers = request.build().unwrap().headers().clone(); - let actual = headers - .get(SESSION_ID_HEADER) + .get("test-my-session-id") .and_then(|value| value.to_str().ok()); - assert_eq!(actual, expected); + assert_eq!(actual, Some("test-session_id-456")); }); } } diff --git a/crates/goose-providers/src/base.rs b/crates/goose-providers/src/base.rs index 855aebd7908f..357215fdaa24 100644 --- a/crates/goose-providers/src/base.rs +++ b/crates/goose-providers/src/base.rs @@ -383,35 +383,22 @@ pub trait Provider: Send + Sync { fn get_name(&self) -> &str; /// Primary streaming method that all providers must implement. - /// - /// Note: Do not add `#[instrument]` here — the call sites (`complete` and - /// `stream_response_from_provider`) create the telemetry span so that - /// `session.id` is set once rather than in every provider. async fn stream( &self, model_config: &ModelConfig, - session_id: &str, system: &str, messages: &[Message], tools: &[Tool], ) -> Result; - /// Complete with a specific model config. - #[tracing::instrument( - skip(self, model_config, session_id, system, messages, tools), - fields(session.id = %session_id, gen_ai.request.model = %model_config.model_name) - )] async fn complete( &self, model_config: &ModelConfig, - session_id: &str, system: &str, messages: &[Message], tools: &[Tool], ) -> Result<(Message, ProviderUsage), ProviderError> { - let stream = self - .stream(model_config, session_id, system, messages, tools) - .await?; + let stream = self.stream(model_config, system, messages, tools).await?; collect_stream(stream).await } diff --git a/crates/goose-providers/src/openai.rs b/crates/goose-providers/src/openai.rs index 9a9719cc6fc5..e44acbb60a04 100644 --- a/crates/goose-providers/src/openai.rs +++ b/crates/goose-providers/src/openai.rs @@ -388,11 +388,7 @@ impl OpenAiProvider { async fn fetch_models_from_api(&self) -> Result, ProviderError> { let models_path = Self::map_base_path(&self.base_path, "models", OPEN_AI_DEFAULT_MODELS_PATH); - let response = self - .api_client - .request(None, &models_path) - .response_get() - .await?; + let response = self.api_client.request(&models_path).response_get().await?; if response.status() == StatusCode::NOT_FOUND { let body = response.text().await.unwrap_or_default(); @@ -427,7 +423,7 @@ impl OpenAiProvider { Self::map_base_path(&self.base_path, "models", OPEN_AI_DEFAULT_MODELS_PATH); let response = self .api_client - .request(None, &models_path) + .request(&models_path) .response_get() .await .ok()?; @@ -581,7 +577,6 @@ impl Provider for OpenAiProvider { async fn stream( &self, model_config: &ModelConfig, - session_id: &str, system: &str, messages: &[Message], tools: &[Tool], @@ -598,7 +593,6 @@ impl Provider for OpenAiProvider { let resp = self .api_client .response_post( - Some(session_id), &Self::map_base_path( &self.base_path, "responses", @@ -659,7 +653,7 @@ impl Provider for OpenAiProvider { .with_retry(|| async { let resp = self .api_client - .response_post(Some(session_id), &self.base_path, &payload) + .response_post(&self.base_path, &payload) .await?; handle_status(resp).await }) diff --git a/crates/goose-providers/src/openai_compatible.rs b/crates/goose-providers/src/openai_compatible.rs index 68aad564a606..5ce39c8173a2 100644 --- a/crates/goose-providers/src/openai_compatible.rs +++ b/crates/goose-providers/src/openai_compatible.rs @@ -78,7 +78,7 @@ impl Provider for OpenAiCompatibleProvider { async fn fetch_supported_models(&self) -> Result, ProviderError> { let response = self .api_client - .response_get(None, "models") + .response_get("models") .await .map_err(|e| ProviderError::RequestFailed(e.to_string()))?; let json = handle_response_openai_compat(response).await?; @@ -105,7 +105,6 @@ impl Provider for OpenAiCompatibleProvider { async fn stream( &self, model_config: &ModelConfig, - session_id: &str, system: &str, messages: &[Message], tools: &[Tool], @@ -124,7 +123,7 @@ impl Provider for OpenAiCompatibleProvider { .with_retry(|| async { let resp = self .api_client - .response_post(Some(session_id), &completions_path, &payload) + .response_post(&completions_path, &payload) .await?; handle_status(resp).await }) diff --git a/crates/goose-server/src/routes/sampling.rs b/crates/goose-server/src/routes/sampling.rs index 63c74b956a3b..7057d61a2835 100644 --- a/crates/goose-server/src/routes/sampling.rs +++ b/crates/goose-server/src/routes/sampling.rs @@ -58,13 +58,15 @@ async fn create_message( tracing::error!("Failed to resolve model config: {}", e); StatusCode::INTERNAL_SERVER_ERROR })?; - let (response, usage) = provider - .complete(&model_config, &session_id, system, &messages, &[]) - .await - .map_err(|e| { - tracing::error!("Sampling completion failed: {}", e); - StatusCode::INTERNAL_SERVER_ERROR - })?; + let (response, usage) = goose::session_context::with_session_id( + Some(session_id.clone()), + provider.complete(&model_config, system, &messages, &[]), + ) + .await + .map_err(|e| { + tracing::error!("Sampling completion failed: {}", e); + StatusCode::INTERNAL_SERVER_ERROR + })?; let text = response.as_concat_text(); diff --git a/crates/goose/examples/databricks_oauth.rs b/crates/goose/examples/databricks_oauth.rs index 7e8a34bb748c..e1883f248e54 100644 --- a/crates/goose/examples/databricks_oauth.rs +++ b/crates/goose/examples/databricks_oauth.rs @@ -19,7 +19,6 @@ async fn main() -> Result<()> { let (response, usage) = provider .complete( &model_config, - "", "You are a helpful assistant.", &[message], &[], diff --git a/crates/goose/examples/image_tool.rs b/crates/goose/examples/image_tool.rs index b6f2949665c5..d2ba3a145621 100644 --- a/crates/goose/examples/image_tool.rs +++ b/crates/goose/examples/image_tool.rs @@ -77,7 +77,6 @@ async fn main() -> Result<()> { let (response, usage) = provider .complete( &model_config, - "", "You are a helpful assistant. Please describe any text you see in the image.", &messages, &[Tool::new("view_image", "View an image", input_schema)], diff --git a/crates/goose/src/acp/provider.rs b/crates/goose/src/acp/provider.rs index 8640c54b0352..86352c5cbeef 100644 --- a/crates/goose/src/acp/provider.rs +++ b/crates/goose/src/acp/provider.rs @@ -465,7 +465,6 @@ impl Provider for AcpProvider { async fn stream( &self, model_config: &ModelConfig, - _session_id: &str, _system: &str, messages: &[Message], _tools: &[Tool], @@ -1736,9 +1735,7 @@ mod tests { Message::user().with_text("current request"), ]; - let result = provider - .stream(&model, "goose-session", "", &messages, &[]) - .await; + let result = provider.stream(&model, "", &messages, &[]).await; assert!(matches!(result, Err(ProviderError::RequestFailed(_)))); let next_claim = provider.claim_handoff_context(&messages); diff --git a/crates/goose/src/acp/server.rs b/crates/goose/src/acp/server.rs index 86ca1023bd9d..a9e862fee847 100644 --- a/crates/goose/src/acp/server.rs +++ b/crates/goose/src/acp/server.rs @@ -1506,15 +1506,16 @@ impl GooseAcpAgent { // common cases without paying for the regular model. let mut llm_outcome: Option = None; for attempt in 0..2 { - match provider - .complete( + match crate::session_context::with_session_id( + Some(sid.0.to_string()), + provider.complete( &fast_model_config, - &sid.0, system, std::slice::from_ref(&message), &[], - ) - .await + ), + ) + .await { Ok((response, _)) => { let summary: String = response @@ -1796,15 +1797,16 @@ impl GooseAcpAgent { // momentarily flaky, without escalating to the regular model. let mut summary: Option = None; for attempt in 0..2 { - match provider - .complete( + match crate::session_context::with_session_id( + Some(sid.0.to_string()), + provider.complete( &fast_model_config, - &sid.0, system, std::slice::from_ref(&message), &[], - ) - .await + ), + ) + .await { Ok((response, _)) => { let s = response diff --git a/crates/goose/src/agents/agent.rs b/crates/goose/src/agents/agent.rs index de29dfedc23d..741b06aa9926 100644 --- a/crates/goose/src/agents/agent.rs +++ b/crates/goose/src/agents/agent.rs @@ -3088,28 +3088,21 @@ impl Agent { ); tracing::info!("Calling provider to generate recipe content"); - let (result, _usage) = self - .provider - .lock() - .await - .as_ref() - .ok_or_else(|| { - let error = anyhow!("Provider not available during recipe creation"); - tracing::error!("{}", error); - error - })? - .complete( - &model_config, - session_id, - &system_prompt, - messages.messages(), - &tools, - ) - .await - .map_err(|e| { - tracing::error!("Provider completion failed during recipe creation: {}", e); - e - })?; + let provider = self.provider.lock().await; + let provider = provider.as_ref().ok_or_else(|| { + let error = anyhow!("Provider not available during recipe creation"); + tracing::error!("{}", error); + error + })?; + let (result, _usage) = crate::session_context::with_session_id( + Some(session_id.to_string()), + provider.complete(&model_config, &system_prompt, messages.messages(), &tools), + ) + .await + .map_err(|e| { + tracing::error!("Provider completion failed during recipe creation: {}", e); + e + })?; let content = result.as_concat_text(); tracing::debug!( @@ -3318,7 +3311,6 @@ mod tests { &self, _: &goose_providers::model::ModelConfig, _: &str, - _: &str, _: &[crate::conversation::message::Message], _: &[rmcp::model::Tool], ) -> Result { @@ -3511,7 +3503,6 @@ exit 0 async fn stream( &self, _model_config: &goose_providers::model::ModelConfig, - _session_id: &str, _system_prompt: &str, _messages: &[Message], _tools: &[Tool], @@ -3534,7 +3525,6 @@ exit 0 async fn stream( &self, _model_config: &goose_providers::model::ModelConfig, - _session_id: &str, _system_prompt: &str, _messages: &[Message], _tools: &[Tool], @@ -3563,7 +3553,6 @@ exit 0 async fn stream( &self, _model_config: &goose_providers::model::ModelConfig, - _session_id: &str, _system_prompt: &str, _messages: &[Message], _tools: &[Tool], diff --git a/crates/goose/src/agents/mcp_client.rs b/crates/goose/src/agents/mcp_client.rs index 919ca3fa4e2f..5b3e27f9dcc0 100644 --- a/crates/goose/src/agents/mcp_client.rs +++ b/crates/goose/src/agents/mcp_client.rs @@ -421,22 +421,18 @@ impl ClientHandler for GooseClient { Some(Value::from(e.to_string())), ) })?; - let (response, usage) = provider - .complete( - &model_config, - session_id.as_deref().unwrap_or(""), - system_prompt, - &provider_ready_messages, - &[], + let (response, usage) = crate::session_context::with_session_id( + session_id.clone(), + provider.complete(&model_config, system_prompt, &provider_ready_messages, &[]), + ) + .await + .map_err(|e| { + ErrorData::new( + ErrorCode::INTERNAL_ERROR, + "Unexpected error while completing the prompt", + Some(Value::from(e.to_string())), ) - .await - .map_err(|e| { - ErrorData::new( - ErrorCode::INTERNAL_ERROR, - "Unexpected error while completing the prompt", - Some(Value::from(e.to_string())), - ) - })?; + })?; Ok(CreateMessageResult::new( SamplingMessage::new( diff --git a/crates/goose/src/agents/platform_extensions/apps.rs b/crates/goose/src/agents/platform_extensions/apps.rs index ee18b586df4c..0edfcf4974b1 100644 --- a/crates/goose/src/agents/platform_extensions/apps.rs +++ b/crates/goose/src/agents/platform_extensions/apps.rs @@ -274,10 +274,12 @@ impl AppsManagerClient { let model_config = self.context.model_config_for_session(session_id).await?; - let (response, usage) = provider - .complete(&model_config, session_id, &system_prompt, &messages, &tools) - .await - .map_err(|e| format!("LLM call failed: {}", e))?; + let (response, usage) = crate::session_context::with_session_id( + Some(session_id.to_string()), + provider.complete(&model_config, &system_prompt, &messages, &tools), + ) + .await + .map_err(|e| format!("LLM call failed: {}", e))?; if let (Some(output), Some(max)) = (usage.usage.output_tokens, model_config.max_tokens) { if output >= max { @@ -317,10 +319,12 @@ impl AppsManagerClient { let model_config = self.context.model_config_for_session(session_id).await?; - let (response, usage) = provider - .complete(&model_config, session_id, &system_prompt, &messages, &tools) - .await - .map_err(|e| format!("LLM call failed: {}", e))?; + let (response, usage) = crate::session_context::with_session_id( + Some(session_id.to_string()), + provider.complete(&model_config, &system_prompt, &messages, &tools), + ) + .await + .map_err(|e| format!("LLM call failed: {}", e))?; if let (Some(output), Some(max)) = (usage.usage.output_tokens, model_config.max_tokens) { if output >= max { diff --git a/crates/goose/src/agents/platform_extensions/summarize.rs b/crates/goose/src/agents/platform_extensions/summarize.rs index fb405b93a46d..d0fc86b74140 100644 --- a/crates/goose/src/agents/platform_extensions/summarize.rs +++ b/crates/goose/src/agents/platform_extensions/summarize.rs @@ -197,10 +197,12 @@ async fn execute_summarize( let user_message = Message::user().with_text(&prompt); - let (response, _usage) = provider - .complete(&model_config, session_id, system, &[user_message], &[]) - .await - .map_err(|e| format!("LLM call failed: {}", e))?; + let (response, _usage) = crate::session_context::with_session_id( + Some(session_id.to_string()), + provider.complete(&model_config, system, &[user_message], &[]), + ) + .await + .map_err(|e| format!("LLM call failed: {}", e))?; let response_text = response .content diff --git a/crates/goose/src/agents/reply_parts.rs b/crates/goose/src/agents/reply_parts.rs index 3e7f6f816b60..c1c366888c02 100644 --- a/crates/goose/src/agents/reply_parts.rs +++ b/crates/goose/src/agents/reply_parts.rs @@ -288,15 +288,16 @@ impl Agent { let model_config = model_config.with_default_thinking_effort(Config::global().get_goose_thinking_effort()); debug!("WAITING_LLM_STREAM_START"); - let stream_result = provider - .stream( + let stream_result = crate::session_context::with_session_id( + Some(session_id.to_string()), + provider.stream( &model_config, - session_id, system_prompt.as_str(), messages_for_provider.messages(), &tools, - ) - .await; + ), + ) + .await; debug!("WAITING_LLM_STREAM_END"); // If there was an error creating the stream, return a stream that yields that error @@ -635,7 +636,6 @@ mod tests { async fn stream( &self, _model_config: &ModelConfig, - _session_id: &str, _system: &str, _messages: &[Message], _tools: &[Tool], diff --git a/crates/goose/src/config/signup_nanogpt/mod.rs b/crates/goose/src/config/signup_nanogpt/mod.rs index ecd6e906ae71..b9b246076459 100644 --- a/crates/goose/src/config/signup_nanogpt/mod.rs +++ b/crates/goose/src/config/signup_nanogpt/mod.rs @@ -38,7 +38,7 @@ async fn poll_for_token(client: &ApiClient, device_code: &str) -> Result let body = json!({ "device_code": device_code }); - let response = client.response_post(None, "poll", &body).await?; + let response = client.response_post("poll", &body).await?; // https://docs.nano-gpt.com/integrations/cli-login#response-codes match response.status().as_u16() { 200 => { @@ -78,7 +78,7 @@ pub async fn complete_nanogpt_auth() -> Result { let client = build_client()?; let body = json!({ "client_name": "goose" }); - let response = client.response_post(None, "start", &body).await?; + let response = client.response_post("start", &body).await?; if !response.status().is_success() { let status = response.status(); diff --git a/crates/goose/src/context_mgmt/mod.rs b/crates/goose/src/context_mgmt/mod.rs index 3553d4fc66cc..59f39a2f5ffc 100644 --- a/crates/goose/src/context_mgmt/mod.rs +++ b/crates/goose/src/context_mgmt/mod.rs @@ -665,7 +665,6 @@ mod tests { async fn stream( &self, _model_config: &ModelConfig, - _session_id: &str, _system: &str, messages: &[Message], _tools: &[Tool], diff --git a/crates/goose/src/dictation/providers.rs b/crates/goose/src/dictation/providers.rs index 8dd3d1ee5307..e62b3be23fa3 100644 --- a/crates/goose/src/dictation/providers.rs +++ b/crates/goose/src/dictation/providers.rs @@ -287,7 +287,7 @@ pub async fn transcribe_with_provider( .text(model_param, model_value); let response = client - .request(None, &endpoint_path) + .request(&endpoint_path) .multipart_post(form) .await .map_err(|e| { diff --git a/crates/goose/src/doctor.rs b/crates/goose/src/doctor.rs index 88a13528d02b..7013d2b46694 100644 --- a/crates/goose/src/doctor.rs +++ b/crates/goose/src/doctor.rs @@ -153,15 +153,16 @@ async fn test_provider( model_config: &goose_providers::model::ModelConfig, ) -> Result<(), ProviderError> { let messages = vec![Message::user().with_text("Say 'hello' and nothing else.")]; - provider - .complete( + crate::session_context::with_session_id( + Some("doctor-check".to_string()), + provider.complete( model_config, - "doctor-check", "Respond as briefly as possible.", &messages, &[], - ) - .await?; + ), + ) + .await?; Ok(()) } diff --git a/crates/goose/src/execution/manager.rs b/crates/goose/src/execution/manager.rs index b2bb1f7bf279..4eea38862802 100644 --- a/crates/goose/src/execution/manager.rs +++ b/crates/goose/src/execution/manager.rs @@ -638,7 +638,6 @@ mod tests { async fn stream( &self, _model_config: &ModelConfig, - _session_id: &str, _system: &str, _messages: &[Message], _tools: &[Tool], diff --git a/crates/goose/src/model_config.rs b/crates/goose/src/model_config.rs index a66df05ef688..cc9a275f5d2b 100644 --- a/crates/goose/src/model_config.rs +++ b/crates/goose/src/model_config.rs @@ -114,9 +114,11 @@ pub async fn complete_fast( .map_err(|e| ProviderError::ExecutionError(e.to_string()))? .with_thinking_effort(ThinkingEffort::Off); - match provider - .complete(&fast_model_config, session_id, system, messages, tools) - .await + match crate::session_context::with_session_id( + Some(session_id.to_string()), + provider.complete(&fast_model_config, system, messages, tools), + ) + .await { Ok(response) => Ok(response), Err(e) if fast_model_config.model_name != model_config.model_name => { @@ -129,9 +131,11 @@ pub async fn complete_fast( let fallback_config = model_config .clone() .with_thinking_effort(ThinkingEffort::Off); - provider - .complete(&fallback_config, session_id, system, messages, tools) - .await + crate::session_context::with_session_id( + Some(session_id.to_string()), + provider.complete(&fallback_config, system, messages, tools), + ) + .await } Err(e) => Err(e), } diff --git a/crates/goose/src/permission/permission_judge.rs b/crates/goose/src/permission/permission_judge.rs index 59db44f1dc62..237266bdaa90 100644 --- a/crates/goose/src/permission/permission_judge.rs +++ b/crates/goose/src/permission/permission_judge.rs @@ -164,15 +164,16 @@ pub async fn detect_read_only_tools( return vec![]; } }; - let res = provider - .complete( + let res = crate::session_context::with_session_id( + Some(session_id.to_string()), + provider.complete( &model_config, - session_id, &system_prompt, check_messages.messages(), std::slice::from_ref(&tool), - ) - .await; + ), + ) + .await; // Process the response and return an empty vector if the response is invalid if let Ok((message, _usage)) = res { diff --git a/crates/goose/src/providers/anthropic_def.rs b/crates/goose/src/providers/anthropic_def.rs index 941e6c93a3f7..b0e695365cfe 100644 --- a/crates/goose/src/providers/anthropic_def.rs +++ b/crates/goose/src/providers/anthropic_def.rs @@ -43,6 +43,7 @@ async fn from_env( }; let api_client = ApiClient::new_with_tls(host, auth, tls_config)? + .with_request_builder(crate::session_context::session_id_request_builder()) .with_header("anthropic-version", ANTHROPIC_API_VERSION)?; Ok(AnthropicProviderBuilder::new(api_client).build()) @@ -85,6 +86,7 @@ pub fn from_custom_config( let format_options = format_options_for_provider(config.preserves_thinking); let mut api_client = ApiClient::new_with_tls(config.base_url, auth, tls_config)? + .with_request_builder(crate::session_context::session_id_request_builder()) .with_header("anthropic-version", ANTHROPIC_API_VERSION)?; if let Some(headers) = &config.headers { diff --git a/crates/goose/src/providers/avian.rs b/crates/goose/src/providers/avian.rs index 05105a1c24bf..a657aa973869 100644 --- a/crates/goose/src/providers/avian.rs +++ b/crates/goose/src/providers/avian.rs @@ -49,7 +49,8 @@ impl ProviderDef for AvianProvider { .unwrap_or_else(|_| AVIAN_API_HOST.to_string()); let api_client = - ApiClient::new_with_tls(host, AuthMethod::BearerToken(api_key), tls_config)?; + ApiClient::new_with_tls(host, AuthMethod::BearerToken(api_key), tls_config)? + .with_request_builder(crate::session_context::session_id_request_builder()); Ok(OpenAiCompatibleProvider::new( AVIAN_PROVIDER_NAME.to_string(), diff --git a/crates/goose/src/providers/azure.rs b/crates/goose/src/providers/azure.rs index f0f10ec9466f..9fd59f2bad78 100644 --- a/crates/goose/src/providers/azure.rs +++ b/crates/goose/src/providers/azure.rs @@ -110,7 +110,8 @@ impl ProviderDef for AzureProvider { host, AuthMethod::Custom(Box::new(auth_provider)), tls_config, - )?; + )? + .with_request_builder(crate::session_context::session_id_request_builder()); if let Some(version) = api_version { api_client = api_client.with_query(vec![("api-version".to_string(), version)]); } diff --git a/crates/goose/src/providers/bedrock.rs b/crates/goose/src/providers/bedrock.rs index b23bb89b90d3..1644b2e38e61 100644 --- a/crates/goose/src/providers/bedrock.rs +++ b/crates/goose/src/providers/bedrock.rs @@ -734,15 +734,15 @@ impl Provider for BedrockProvider { async fn stream( &self, model_config: &ModelConfig, - session_id: &str, system: &str, messages: &[Message], tools: &[Tool], ) -> Result { + let session_id = crate::session_context::current_session_id().unwrap_or_default(); let session_id_opt = if session_id.is_empty() { None } else { - Some(session_id) + Some(session_id.as_str()) }; let without_prefix = model_config @@ -1102,7 +1102,7 @@ mod tests { let messages = vec![crate::conversation::message::Message::user().with_text("hi")]; let mut stream = provider - .stream(&model.clone(), "", "", &messages, &[]) + .stream(&model.clone(), "", &messages, &[]) .await .unwrap(); diff --git a/crates/goose/src/providers/chatgpt_codex.rs b/crates/goose/src/providers/chatgpt_codex.rs index cb74d8837a7c..b80f5fc9a129 100644 --- a/crates/goose/src/providers/chatgpt_codex.rs +++ b/crates/goose/src/providers/chatgpt_codex.rs @@ -1,10 +1,9 @@ use crate::config::paths::Paths; use crate::conversation::message::{Message, MessageContent}; -use crate::providers::api_client::AuthProvider; +use crate::providers::api_client::{AuthProvider, RequestBuilderDecorator}; use crate::providers::base::{ConfigKey, MessageStream, Provider, ProviderDef, ProviderMetadata}; use crate::providers::openai_compatible::handle_status; use crate::providers::retry::ProviderRetry; -use crate::session_context::SESSION_ID_HEADER; use anyhow::{anyhow, Result}; use async_stream::try_stream; use async_trait::async_trait; @@ -18,7 +17,6 @@ use goose_providers::formats::openai_responses::responses_api_to_streaming_messa use goose_providers::model::ModelConfig; use jsonwebtoken::jwk::JwkSet; use jsonwebtoken::{decode, decode_header, DecodingKey, Validation}; -use reqwest::header::{HeaderName, HeaderValue}; use rmcp::model::{RawContent, Role, Tool}; use serde::{Deserialize, Serialize}; use serde_json::{json, Value}; @@ -875,12 +873,14 @@ impl AuthProvider for ChatGptCodexAuthProvider { } } -#[derive(Debug, serde::Serialize)] +#[derive(serde::Serialize)] pub struct ChatGptCodexProvider { #[serde(skip)] auth_provider: Arc, #[serde(skip)] name: String, + #[serde(skip)] + request_builder: RequestBuilderDecorator, } impl ChatGptCodexProvider { @@ -899,14 +899,11 @@ impl ChatGptCodexProvider { Ok(Self { auth_provider, name: CHATGPT_CODEX_PROVIDER_NAME.to_string(), + request_builder: crate::session_context::session_id_request_builder(), }) } - async fn post_streaming( - &self, - session_id: Option<&str>, - payload: &Value, - ) -> Result { + async fn post_streaming(&self, payload: &Value) -> Result { let token_data = self .auth_provider .get_valid_token() @@ -922,16 +919,8 @@ impl ChatGptCodexProvider { ); } - if let Some(session_id) = session_id.filter(|id| !id.is_empty()) { - headers.insert( - HeaderName::from_static(SESSION_ID_HEADER), - HeaderValue::from_str(session_id) - .map_err(|e| ProviderError::ExecutionError(e.to_string()))?, - ); - } - let client = reqwest::Client::new(); - let response = client + let request = client .post(format!("{}/responses", CODEX_API_ENDPOINT)) .header( "Authorization", @@ -939,7 +928,10 @@ impl ChatGptCodexProvider { ) .header("Content-Type", "application/json") .headers(headers) - .json(payload) + .json(payload); + + let response = (self.request_builder)(request) + .map_err(|e| ProviderError::ExecutionError(e.to_string()))? .send() .await .map_err(|e| ProviderError::RequestFailed(e.to_string()))?; @@ -988,7 +980,6 @@ impl Provider for ChatGptCodexProvider { async fn stream( &self, model_config: &ModelConfig, - session_id: &str, system: &str, messages: &[Message], tools: &[Tool], @@ -1000,7 +991,7 @@ impl Provider for ChatGptCodexProvider { let response = self .with_retry(|| async { let payload_clone = payload.clone(); - self.post_streaming(Some(session_id), &payload_clone).await + self.post_streaming(&payload_clone).await }) .await?; diff --git a/crates/goose/src/providers/claude_code.rs b/crates/goose/src/providers/claude_code.rs index 4adbc867f7f6..c2a99b5955f8 100644 --- a/crates/goose/src/providers/claude_code.rs +++ b/crates/goose/src/providers/claude_code.rs @@ -714,11 +714,11 @@ impl Provider for ClaudeCodeProvider { async fn stream( &self, model_config: &ModelConfig, - session_id: &str, system: &str, messages: &[Message], _tools: &[Tool], ) -> Result { + let session_id = crate::session_context::current_session_id().unwrap_or_default(); if super::cli_common::is_session_description_request(system) { let (message, usage) = super::cli_common::generate_simple_session_description( &model_config.model_name, @@ -735,7 +735,7 @@ impl Provider for ClaudeCodeProvider { // Prepare the payload outside the lock — these don't need the process. let blocks = self.last_user_content_blocks(messages); - let ndjson_line = build_stream_json_input(&blocks, session_id); + let ndjson_line = build_stream_json_input(&blocks, &session_id); let model_name = model_config.model_name.clone(); let message_id = uuid::Uuid::new_v4().to_string(); let pending_confirmations = Arc::clone(&self.pending_confirmations); @@ -1316,10 +1316,7 @@ mod tests { let messages = vec![Message::user().with_text("test")]; let model = ModelConfig::new(CLAUDE_CODE_DEFAULT_MODEL) .with_canonical_limits(CLAUDE_CODE_PROVIDER_NAME); - let stream = provider - .stream(&model, "test-session", "", &messages, &[]) - .await - .unwrap(); + let stream = provider.stream(&model, "", &messages, &[]).await.unwrap(); (provider, stream, stdin_reader) } @@ -1526,10 +1523,7 @@ mod tests { let messages = vec![Message::user().with_text("test")]; let model = ModelConfig::new(CLAUDE_CODE_DEFAULT_MODEL) .with_canonical_limits(CLAUDE_CODE_PROVIDER_NAME); - let mut stream = provider - .stream(&model, "test-session", "", &messages, &[]) - .await - .unwrap(); + let mut stream = provider.stream(&model, "", &messages, &[]).await.unwrap(); while let Some(item) = stream.next().await { item.unwrap(); diff --git a/crates/goose/src/providers/codex.rs b/crates/goose/src/providers/codex.rs index c3722640ca8a..44ca5ddf48b6 100644 --- a/crates/goose/src/providers/codex.rs +++ b/crates/goose/src/providers/codex.rs @@ -683,11 +683,11 @@ impl Provider for CodexProvider { async fn stream( &self, model_config: &ModelConfig, - session_id: &str, system: &str, messages: &[Message], tools: &[Tool], ) -> Result { + let session_id = crate::session_context::current_session_id().unwrap_or_default(); if super::cli_common::is_session_description_request(system) { let (message, provider_usage) = super::cli_common::generate_simple_session_description( &model_config.model_name, @@ -701,7 +701,7 @@ impl Provider for CodexProvider { let goose_mode = { let map = self.mode_by_session.read().await; - map.get(session_id).copied().unwrap_or_default() + map.get(&session_id).copied().unwrap_or_default() }; let lines = self .execute_command(model_config, system, messages, tools, goose_mode) diff --git a/crates/goose/src/providers/cursor_agent.rs b/crates/goose/src/providers/cursor_agent.rs index b36440fcff49..22c5186ef2e9 100644 --- a/crates/goose/src/providers/cursor_agent.rs +++ b/crates/goose/src/providers/cursor_agent.rs @@ -324,7 +324,6 @@ impl Provider for CursorAgentProvider { async fn stream( &self, model_config: &ModelConfig, - _session_id: &str, // CLI has no external session-id flag to propagate. system: &str, messages: &[Message], tools: &[Tool], diff --git a/crates/goose/src/providers/databricks.rs b/crates/goose/src/providers/databricks.rs index b9d1c0126185..b3e9ce7dd049 100644 --- a/crates/goose/src/providers/databricks.rs +++ b/crates/goose/src/providers/databricks.rs @@ -137,7 +137,8 @@ impl DatabricksProvider { auth_method, Duration::from_secs(DEFAULT_PROVIDER_TIMEOUT_SECS), tls_config.clone(), - )?; + )? + .with_request_builder(crate::session_context::session_id_request_builder()); Ok(Self { api_client, @@ -335,13 +336,10 @@ impl DatabricksProvider { ) -> Result { let response = self .api_client - .request( - None, - &format!( - "api/2.0/serving-endpoints/{}", - urlencoding::encode(endpoint_name) - ), - ) + .request(&format!( + "api/2.0/serving-endpoints/{}", + urlencoding::encode(endpoint_name) + )) .response_get() .await .map_err(|e| { @@ -565,11 +563,11 @@ impl Provider for DatabricksProvider { async fn stream( &self, model_config: &ModelConfig, - session_id: &str, system: &str, messages: &[Message], tools: &[Tool], ) -> Result { + let session_id = crate::session_context::current_session_id().unwrap_or_default(); let (endpoint_name, _) = extract_reasoning_effort(&model_config.model_name); let endpoint_info = self.resolve_endpoint_info_cached(&endpoint_name).await.ok(); let effective_model_name = endpoint_info @@ -585,7 +583,7 @@ impl Provider for DatabricksProvider { } else { self.get_endpoint_path(&model_config.model_name, is_responses_model) }; - let client_request_id = self.build_client_request_id(session_id); + let client_request_id = self.build_client_request_id(&session_id); if is_responses_model { let responses_model_config; @@ -625,10 +623,7 @@ impl Provider for DatabricksProvider { let response = self .with_retry(|| async { let payload_clone = payload.clone(); - let resp = self - .api_client - .response_post(Some(session_id), &path, &payload_clone) - .await?; + let resp = self.api_client.response_post(&path, &payload_clone).await?; handle_status(resp).await }) .await @@ -688,10 +683,7 @@ impl Provider for DatabricksProvider { let mut log = start_log(model_config, &payload)?; let response = self .with_retry(|| async { - let resp = self - .api_client - .response_post(Some(session_id), &path, &payload) - .await?; + let resp = self.api_client.response_post(&path, &payload).await?; if !resp.status().is_success() { let status = resp.status(); let url = sanitize_url(resp.url().as_str()); @@ -708,10 +700,7 @@ impl Provider for DatabricksProvider { Err(e) if e.to_string().contains("stream_options") => { payload.as_object_mut().unwrap().remove("stream_options"); self.with_retry(|| async { - let resp = self - .api_client - .response_post(Some(session_id), &path, &payload) - .await?; + let resp = self.api_client.response_post(&path, &payload).await?; if !resp.status().is_success() { let status = resp.status(); let url = sanitize_url(resp.url().as_str()); @@ -753,7 +742,7 @@ impl Provider for DatabricksProvider { async fn fetch_supported_model_info(&self) -> Result, ProviderError> { let response = self .api_client - .request(None, "api/2.0/serving-endpoints") + .request("api/2.0/serving-endpoints") .response_get() .await .map_err(|e| { diff --git a/crates/goose/src/providers/databricks_v2.rs b/crates/goose/src/providers/databricks_v2.rs index 560d0eb7747a..fb289b13a0e9 100644 --- a/crates/goose/src/providers/databricks_v2.rs +++ b/crates/goose/src/providers/databricks_v2.rs @@ -118,7 +118,8 @@ impl DatabricksV2Provider { auth_method, Duration::from_secs(DEFAULT_PROVIDER_TIMEOUT_SECS), tls_config, - )?; + )? + .with_request_builder(crate::session_context::session_id_request_builder()); Ok(Self { api_client, @@ -217,7 +218,6 @@ impl DatabricksV2Provider { async fn stream_openai_responses( &self, model_config: &ModelConfig, - session_id: &str, system: &str, messages: &[Message], tools: &[Tool], @@ -231,7 +231,7 @@ impl DatabricksV2Provider { .with_retry(|| async { let resp = self .api_client - .response_post(Some(session_id), "ai-gateway/openai/v1/responses", &payload) + .response_post("ai-gateway/openai/v1/responses", &payload) .await?; handle_status(resp).await }) @@ -246,7 +246,6 @@ impl DatabricksV2Provider { async fn stream_mlflow_chat_completions( &self, model_config: &ModelConfig, - session_id: &str, system: &str, messages: &[Message], tools: &[Tool], @@ -268,11 +267,7 @@ impl DatabricksV2Provider { .with_retry(|| async { let resp = self .api_client - .response_post( - Some(session_id), - "ai-gateway/mlflow/v1/chat/completions", - &payload, - ) + .response_post("ai-gateway/mlflow/v1/chat/completions", &payload) .await?; handle_status(resp).await }) @@ -287,7 +282,6 @@ impl DatabricksV2Provider { async fn stream_anthropic_messages( &self, model_config: &ModelConfig, - session_id: &str, system: &str, messages: &[Message], tools: &[Tool], @@ -307,11 +301,7 @@ impl DatabricksV2Provider { .with_retry(|| async { let resp = self .api_client - .response_post( - Some(session_id), - "ai-gateway/anthropic/v1/messages", - &payload, - ) + .response_post("ai-gateway/anthropic/v1/messages", &payload) .await?; handle_status(resp).await }) @@ -385,29 +375,22 @@ impl Provider for DatabricksV2Provider { async fn stream( &self, model_config: &ModelConfig, - session_id: &str, system: &str, messages: &[Message], tools: &[Tool], ) -> Result { match Self::route_for_model(&model_config.model_name) { DatabricksV2Route::OpenAiResponses => { - self.stream_openai_responses(model_config, session_id, system, messages, tools) + self.stream_openai_responses(model_config, system, messages, tools) .await } DatabricksV2Route::AnthropicMessages => { - self.stream_anthropic_messages(model_config, session_id, system, messages, tools) + self.stream_anthropic_messages(model_config, system, messages, tools) .await } DatabricksV2Route::MlflowChatCompletions => { - self.stream_mlflow_chat_completions( - model_config, - session_id, - system, - messages, - tools, - ) - .await + self.stream_mlflow_chat_completions(model_config, system, messages, tools) + .await } } } @@ -425,15 +408,11 @@ impl Provider for DatabricksV2Provider { path.push_str(&format!("&page_token={}", urlencoding::encode(token))); } - let response = self - .api_client - .response_get(None, &path) - .await - .map_err(|e| { - ProviderError::RequestFailed(format!( - "Failed to fetch Databricks AI Gateway endpoints: {e}" - )) - })?; + let response = self.api_client.response_get(&path).await.map_err(|e| { + ProviderError::RequestFailed(format!( + "Failed to fetch Databricks AI Gateway endpoints: {e}" + )) + })?; if !response.status().is_success() { let status = response.status(); diff --git a/crates/goose/src/providers/gcpvertexai.rs b/crates/goose/src/providers/gcpvertexai.rs index 0b6240f0ff64..1230d41dc852 100644 --- a/crates/goose/src/providers/gcpvertexai.rs +++ b/crates/goose/src/providers/gcpvertexai.rs @@ -15,6 +15,7 @@ use tokio_util::io::StreamReader; use url::Url; use crate::conversation::message::Message; +use crate::providers::api_client::RequestBuilderDecorator; use crate::providers::base::{ ConfigKey, MessageStream, Provider, ProviderDef, ProviderMetadata, DEFAULT_PROVIDER_TIMEOUT_SECS, @@ -28,7 +29,6 @@ use crate::providers::formats::gcpvertexai::{ use crate::providers::gcpauth::GcpAuth; use crate::providers::openai_compatible::{map_http_error_to_provider_error, sanitize_url}; use crate::providers::retry::RetryConfig; -use crate::session_context::SESSION_ID_HEADER; use goose_providers::errors::ProviderError; use goose_providers::request_log::{start_log, LoggerHandleExt}; use rmcp::model::Tool; @@ -134,7 +134,7 @@ enum GcpVertexAIError { /// This provider enables interaction with various AI models hosted on GCP Vertex AI, /// including Claude and Gemini model families. It handles authentication, request routing, /// and response processing for the Vertex AI API endpoints. -#[derive(Debug, serde::Serialize)] +#[derive(serde::Serialize)] pub struct GcpVertexAIProvider { /// HTTP client for making API requests #[serde(skip)] @@ -153,6 +153,8 @@ pub struct GcpVertexAIProvider { retry_config: RetryConfig, #[serde(skip)] name: String, + #[serde(skip)] + request_builder: RequestBuilderDecorator, } impl GcpVertexAIProvider { @@ -188,6 +190,7 @@ impl GcpVertexAIProvider { location, retry_config, name: GCP_VERTEX_AI_PROVIDER_NAME.to_string(), + request_builder: crate::session_context::session_id_request_builder(), }) } @@ -276,7 +279,6 @@ impl GcpVertexAIProvider { async fn send_request_with_retry( &self, - session_id: Option<&str>, url: Url, payload: &Value, ) -> Result { @@ -312,17 +314,14 @@ impl GcpVertexAIProvider { } }; - let mut request = self + let request = self .client .post(url.clone()) .json(payload) .header("Authorization", auth_header); - if let Some(session_id) = session_id.filter(|id| !id.is_empty()) { - request = request.header(SESSION_ID_HEADER, session_id); - } - - let response = request + let response = (self.request_builder)(request) + .map_err(|e| ProviderError::ExecutionError(e.to_string()))? .send() .await .map_err(|e| ProviderError::RequestFailed(e.to_string()))?; @@ -393,7 +392,6 @@ impl GcpVertexAIProvider { async fn post_stream_with_location( &self, model: &ModelConfig, - session_id: Option<&str>, payload: &Value, context: &RequestContext, location: &str, @@ -402,18 +400,17 @@ impl GcpVertexAIProvider { .build_request_url(model, context.provider(), location, true) .map_err(|e| ProviderError::RequestFailed(e.to_string()))?; - self.send_request_with_retry(session_id, url, payload).await + self.send_request_with_retry(url, payload).await } async fn post_stream( &self, model: &ModelConfig, - session_id: Option<&str>, payload: &Value, context: &RequestContext, ) -> Result { let result = self - .post_stream_with_location(model, session_id, payload, context, &self.location) + .post_stream_with_location(model, payload, context, &self.location) .await; if self.location == context.model.known_location().to_string() || result.is_ok() { @@ -430,7 +427,7 @@ impl GcpVertexAIProvider { "Trying known location {known_location} for {model_name} instead of {configured_location}: {msg}" ); - self.post_stream_with_location(model, session_id, payload, context, &known_location) + self.post_stream_with_location(model, payload, context, &known_location) .await } _ => result, @@ -604,7 +601,6 @@ impl Provider for GcpVertexAIProvider { async fn stream( &self, model_config: &ModelConfig, - session_id: &str, system: &str, messages: &[Message], tools: &[Tool], @@ -620,7 +616,7 @@ impl Provider for GcpVertexAIProvider { let mut log = start_log(model_config, &request)?; let response = self - .post_stream(model_config, Some(session_id), &request, &context) + .post_stream(model_config, &request, &context) .await .inspect_err(|e| { let _ = log.error(e); diff --git a/crates/goose/src/providers/gemini_cli.rs b/crates/goose/src/providers/gemini_cli.rs index 665d91081df5..9fd1777a6037 100644 --- a/crates/goose/src/providers/gemini_cli.rs +++ b/crates/goose/src/providers/gemini_cli.rs @@ -204,7 +204,6 @@ impl Provider for GeminiCliProvider { async fn stream( &self, model_config: &ModelConfig, - _session_id: &str, // CLI has no external session-id flag to propagate. system: &str, messages: &[Message], _tools: &[Tool], diff --git a/crates/goose/src/providers/gemini_oauth.rs b/crates/goose/src/providers/gemini_oauth.rs index dfe5ce720de0..a872356d6104 100644 --- a/crates/goose/src/providers/gemini_oauth.rs +++ b/crates/goose/src/providers/gemini_oauth.rs @@ -1,5 +1,6 @@ use crate::config::paths::Paths; use crate::conversation::message::Message; +use crate::providers::api_client::RequestBuilderDecorator; use crate::providers::base::{ ConfigKey, MessageStream, Provider, ProviderDef, ProviderMetadata, DEFAULT_PROVIDER_TIMEOUT_SECS, @@ -13,7 +14,6 @@ use goose_providers::request_log::{start_log, LoggerHandleExt}; const GEMINI_OAUTH_DEFAULT_MODEL: &str = "gemini-3-flash-preview"; const GEMINI_OAUTH_DEFAULT_FAST_MODEL: &str = "gemini-2.5-flash-lite"; use crate::providers::retry::ProviderRetry; -use crate::session_context::SESSION_ID_HEADER; use anyhow::{anyhow, Result}; use async_stream::try_stream; use async_trait::async_trait; @@ -22,7 +22,6 @@ use base64::Engine; use chrono::{DateTime, Utc}; use futures::future::BoxFuture; use futures::TryStreamExt; -use reqwest::header::{HeaderName, HeaderValue}; use rmcp::model::Tool; use serde::{Deserialize, Serialize}; use serde_json::{json, Value}; @@ -827,12 +826,14 @@ fn parse_retry_delay(body: &str) -> Option { // Provider // --------------------------------------------------------------------------- -#[derive(Debug, serde::Serialize)] +#[derive(serde::Serialize)] pub struct GeminiOAuthProvider { #[serde(skip)] token_provider: Arc, #[serde(skip)] name: String, + #[serde(skip)] + request_builder: RequestBuilderDecorator, } impl GeminiOAuthProvider { @@ -846,6 +847,7 @@ impl GeminiOAuthProvider { Ok(Self { token_provider, name: GEMINI_OAUTH_PROVIDER_NAME.to_string(), + request_builder: crate::session_context::session_id_request_builder(), }) } @@ -856,7 +858,6 @@ impl GeminiOAuthProvider { async fn post_stream( &self, - session_id: Option<&str>, model_name: &str, payload: &Value, ) -> Result { @@ -873,7 +874,7 @@ impl GeminiOAuthProvider { CODE_ASSIST_ENDPOINT, CODE_ASSIST_API_VERSION ); - let mut request = HTTP_CLIENT + let request = HTTP_CLIENT .post(&url) .header( "Authorization", @@ -881,14 +882,8 @@ impl GeminiOAuthProvider { ) .header("Content-Type", "application/json"); - if let Some(session_id) = session_id.filter(|id| !id.is_empty()) { - if let Ok(val) = HeaderValue::from_str(session_id) { - request = request.header(HeaderName::from_static(SESSION_ID_HEADER), val); - } - } - - let response = request - .json(&wrapped) + let response = (self.request_builder)(request.json(&wrapped)) + .map_err(|e| ProviderError::ExecutionError(e.to_string()))? .send() .await .map_err(|e| ProviderError::RequestFailed(e.to_string()))?; @@ -982,7 +977,6 @@ impl Provider for GeminiOAuthProvider { async fn stream( &self, model_config: &ModelConfig, - session_id: &str, system: &str, messages: &[Message], tools: &[Tool], @@ -991,10 +985,7 @@ impl Provider for GeminiOAuthProvider { let mut log = start_log(model_config, &payload)?; let response = self - .with_retry(|| async { - self.post_stream(Some(session_id), &model_config.model_name, &payload) - .await - }) + .with_retry(|| async { self.post_stream(&model_config.model_name, &payload).await }) .await .inspect_err(|e| { let _ = log.error(e); diff --git a/crates/goose/src/providers/githubcopilot.rs b/crates/goose/src/providers/githubcopilot.rs index 8c3e9fce60f9..001164635e8e 100644 --- a/crates/goose/src/providers/githubcopilot.rs +++ b/crates/goose/src/providers/githubcopilot.rs @@ -262,7 +262,6 @@ impl GithubCopilotProvider { async fn post( &self, - session_id: Option<&str>, path: &str, is_user_initiated: bool, payload: &mut Value, @@ -277,10 +276,11 @@ impl GithubCopilotProvider { let initiator = if is_user_initiated { "user" } else { "agent" }; headers.insert("X-Initiator", initiator.parse().unwrap()); let api_client = ApiClient::new_with_tls(endpoint.clone(), auth, self.tls_config.clone())? + .with_request_builder(crate::session_context::session_id_request_builder()) .with_headers(headers)?; api_client - .response_post(session_id, path, payload) + .response_post(path, payload) .await .map_err(|e| e.into()) } @@ -397,7 +397,6 @@ impl GithubCopilotProvider { async fn stream_responses( &self, model_config: &ModelConfig, - session_id: &str, is_user_initiated: bool, system: &str, messages: &[Message], @@ -415,7 +414,6 @@ impl GithubCopilotProvider { let mut payload_clone = payload.clone(); let resp = self .post( - Some(session_id), "responses", is_user_initiated, &mut payload_clone, @@ -436,7 +434,6 @@ impl GithubCopilotProvider { async fn stream_chat_completions( &self, model_config: &ModelConfig, - session_id: &str, is_user_initiated: bool, system: &str, messages: &[Message], @@ -463,7 +460,6 @@ impl GithubCopilotProvider { let mut payload_clone = payload.clone(); let resp = self .post( - Some(session_id), "chat/completions", is_user_initiated, &mut payload_clone, @@ -479,11 +475,6 @@ impl GithubCopilotProvider { stream_openai_compat(response, log) } else { - let session_id_opt = if session_id.is_empty() { - None - } else { - Some(session_id) - }; let payload = create_request( model_config, system, @@ -498,7 +489,6 @@ impl GithubCopilotProvider { .with_retry(|| async { let mut payload_clone = payload.clone(); self.post( - session_id_opt, "chat/completions", is_user_initiated, &mut payload_clone, @@ -563,25 +553,16 @@ impl Provider for GithubCopilotProvider { &self.name } - #[tracing::instrument( - skip(self, model_config, session_id, system, messages, tools), - fields(session.id = %session_id, gen_ai.request.model = %model_config.model_name) - )] async fn complete( &self, model_config: &ModelConfig, - session_id: &str, system: &str, messages: &[Message], tools: &[Tool], ) -> Result<(Message, ProviderUsage), ProviderError> { IS_AGENT_CALL .scope(true, async { - collect_stream( - self.stream(model_config, session_id, system, messages, tools) - .await?, - ) - .await + collect_stream(self.stream(model_config, system, messages, tools).await?).await }) .await } @@ -589,7 +570,6 @@ impl Provider for GithubCopilotProvider { async fn stream( &self, model_config: &ModelConfig, - session_id: &str, system: &str, messages: &[Message], tools: &[Tool], @@ -606,7 +586,6 @@ impl Provider for GithubCopilotProvider { if is_openai_responses_model(&model_config.model_name) { self.stream_responses( model_config, - session_id, is_user_initiated, system, messages, @@ -617,7 +596,6 @@ impl Provider for GithubCopilotProvider { } else { self.stream_chat_completions( model_config, - session_id, is_user_initiated, system, messages, diff --git a/crates/goose/src/providers/google.rs b/crates/goose/src/providers/google.rs index 3ab13a819fc2..d00a40f0b03d 100644 --- a/crates/goose/src/providers/google.rs +++ b/crates/goose/src/providers/google.rs @@ -80,6 +80,7 @@ impl GoogleProvider { }; let api_client = ApiClient::new_with_tls(host, auth, tls_config)? + .with_request_builder(crate::session_context::session_id_request_builder()) .with_header("Content-Type", "application/json")?; Ok(Self { @@ -90,15 +91,11 @@ impl GoogleProvider { async fn post_stream( &self, - session_id: Option<&str>, model_name: &str, payload: &Value, ) -> Result { let path = format!("v1beta/models/{}:streamGenerateContent?alt=sse", model_name); - let response = self - .api_client - .response_post(session_id, &path, payload) - .await?; + let response = self.api_client.response_post(&path, payload).await?; handle_status(response).await } } @@ -147,7 +144,7 @@ impl Provider for GoogleProvider { async fn fetch_supported_models(&self) -> Result, ProviderError> { let response = self .api_client - .request(None, "v1beta/models") + .request("v1beta/models") .response_get() .await?; let status = response.status(); @@ -179,7 +176,6 @@ impl Provider for GoogleProvider { async fn stream( &self, model_config: &ModelConfig, - session_id: &str, system: &str, messages: &[Message], tools: &[Tool], @@ -188,10 +184,7 @@ impl Provider for GoogleProvider { let mut log = start_log(model_config, &payload)?; let response = self - .with_retry(|| async { - self.post_stream(Some(session_id), &model_config.model_name, &payload) - .await - }) + .with_retry(|| async { self.post_stream(&model_config.model_name, &payload).await }) .await .inspect_err(|e| { let _ = log.error(e); diff --git a/crates/goose/src/providers/huggingface.rs b/crates/goose/src/providers/huggingface.rs index b24989fdedcc..b1e52d118539 100644 --- a/crates/goose/src/providers/huggingface.rs +++ b/crates/goose/src/providers/huggingface.rs @@ -97,6 +97,7 @@ impl HuggingFaceProvider { std::time::Duration::from_secs(timeout_secs), tls_config, )? + .with_request_builder(crate::session_context::session_id_request_builder()) .with_query(query_params); if let Some(headers) = &config.headers { @@ -158,13 +159,12 @@ impl Provider for HuggingFaceProvider { async fn stream( &self, model_config: &ModelConfig, - session_id: &str, system: &str, messages: &[Message], tools: &[Tool], ) -> Result { self.inner - .stream(model_config, session_id, system, messages, tools) + .stream(model_config, system, messages, tools) .await } } @@ -206,7 +206,8 @@ impl ProviderDef for HuggingFaceProvider { let host: String = config .get_param("HF_HOST") .unwrap_or_else(|_| HUGGINGFACE_API_HOST.to_string()); - let api_client = ApiClient::new_with_tls(host, auth_method, tls_config)?; + let api_client = ApiClient::new_with_tls(host, auth_method, tls_config)? + .with_request_builder(crate::session_context::session_id_request_builder()); Ok(Self { inner: OpenAiCompatibleProvider::new( diff --git a/crates/goose/src/providers/kimicode.rs b/crates/goose/src/providers/kimicode.rs index 378a7d61c2ff..0fc5e1c45daa 100644 --- a/crates/goose/src/providers/kimicode.rs +++ b/crates/goose/src/providers/kimicode.rs @@ -1,6 +1,5 @@ use crate::config::paths::Paths; use crate::config::Config; -use crate::session_context::SESSION_ID_HEADER; use anyhow::Result; use async_stream::try_stream; use async_trait::async_trait; @@ -17,6 +16,7 @@ use tokio::pin; use tokio_util::io::StreamReader; use uuid::Uuid; +use super::api_client::RequestBuilderDecorator; use super::base::{ ConfigKey, MessageStream, Provider, ProviderDef, ProviderMetadata, DEFAULT_PROVIDER_TIMEOUT_SECS, @@ -138,7 +138,7 @@ impl TokenCache { // ── Provider ───────────────────────────────────────────────────────────────── -#[derive(Debug, serde::Serialize)] +#[derive(serde::Serialize)] pub struct KimiCodeProvider { #[serde(skip)] client: Client, @@ -154,6 +154,8 @@ pub struct KimiCodeProvider { api_base: String, #[serde(skip)] name: String, + #[serde(skip)] + request_builder: RequestBuilderDecorator, } impl KimiCodeProvider { @@ -176,6 +178,7 @@ impl KimiCodeProvider { auth_host: KIMI_AUTH_HOST.to_string(), api_base: KIMI_API_BASE.to_string(), name: KIMI_CODE_PROVIDER_NAME.to_string(), + request_builder: crate::session_context::session_id_request_builder(), }) } @@ -306,27 +309,20 @@ impl KimiCodeProvider { // ── HTTP ───────────────────────────────────────────────────────────────── - async fn post( - &self, - session_id: Option<&str>, - payload: &Value, - ) -> Result { + async fn post(&self, payload: &Value) -> Result { let access_token = self.get_access_token().await.map_err(|e| { ProviderError::Authentication(format!("Failed to get Kimi access token: {}", e)) })?; - let mut builder = self + let builder = self .client .post(format!("{}/v1/messages", self.api_base)) .bearer_auth(access_token) .headers(self.kimi_headers()) .json(payload); - if let Some(sid) = session_id { - builder = builder.header(SESSION_ID_HEADER, sid); - } - - builder + (self.request_builder)(builder) + .map_err(|e| ProviderError::ExecutionError(e.to_string()))? .send() .await .map_err(|e| ProviderError::RequestFailed(e.to_string())) @@ -385,7 +381,6 @@ impl Provider for KimiCodeProvider { async fn stream( &self, model_config: &ModelConfig, - session_id: &str, system: &str, messages: &[Message], tools: &[Tool], @@ -409,7 +404,7 @@ impl Provider for KimiCodeProvider { let response = self .with_retry(|| async { - let resp = self.post(Some(session_id), &payload).await?; + let resp = self.post(&payload).await?; handle_status(resp).await }) .await @@ -508,6 +503,7 @@ mod tests { auth_host: server_uri.to_string(), api_base: server_uri.to_string(), name: KIMI_CODE_PROVIDER_NAME.to_string(), + request_builder: std::sync::Arc::new(Ok), } } diff --git a/crates/goose/src/providers/litellm.rs b/crates/goose/src/providers/litellm.rs index 9fed2f93e2ac..4f3d8525f70b 100644 --- a/crates/goose/src/providers/litellm.rs +++ b/crates/goose/src/providers/litellm.rs @@ -69,7 +69,8 @@ impl LiteLLMProvider { auth, std::time::Duration::from_secs(timeout_secs), tls_config, - )?; + )? + .with_request_builder(crate::session_context::session_id_request_builder()); if let Some(headers) = custom_headers { let mut header_map = reqwest::header::HeaderMap::new(); @@ -97,11 +98,7 @@ impl LiteLLMProvider { } async fn fetch_models_from_api(&self) -> Result, ProviderError> { - let response = self - .api_client - .request(None, "model/info") - .response_get() - .await?; + let response = self.api_client.request("model/info").response_get().await?; if !response.status().is_success() { return Err(ProviderError::RequestFailed(format!( @@ -139,14 +136,10 @@ impl LiteLLMProvider { Ok(models) } - async fn post( - &self, - session_id: Option<&str>, - payload: &Value, - ) -> Result { + async fn post(&self, payload: &Value) -> Result { let response = self .api_client - .response_post(session_id, &self.base_path, payload) + .response_post(&self.base_path, payload) .await?; handle_response_openai_compat(response).await } @@ -235,16 +228,10 @@ impl Provider for LiteLLMProvider { async fn stream( &self, model_config: &ModelConfig, - session_id: &str, system: &str, messages: &[Message], tools: &[Tool], ) -> Result { - let session_id = if session_id.is_empty() { - None - } else { - Some(session_id) - }; let mut payload = goose_providers::formats::openai::create_request( model_config, system, @@ -261,7 +248,7 @@ impl Provider for LiteLLMProvider { let response = self .with_retry(|| async { let payload_clone = payload.clone(); - self.post(session_id, &payload_clone).await + self.post(&payload_clone).await }) .await?; diff --git a/crates/goose/src/providers/local_inference.rs b/crates/goose/src/providers/local_inference.rs index fc33f06d0b9f..504a17a5aed1 100644 --- a/crates/goose/src/providers/local_inference.rs +++ b/crates/goose/src/providers/local_inference.rs @@ -563,7 +563,6 @@ impl Provider for LocalInferenceProvider { async fn stream( &self, model_config: &ModelConfig, - _session_id: &str, system: &str, messages: &[Message], tools: &[Tool], diff --git a/crates/goose/src/providers/nanogpt.rs b/crates/goose/src/providers/nanogpt.rs index 03cee2b308b9..56829b472916 100644 --- a/crates/goose/src/providers/nanogpt.rs +++ b/crates/goose/src/providers/nanogpt.rs @@ -51,7 +51,7 @@ impl NanoGptProvider { Err(_) => return false, }; - match client.response_get(None, "usage").await { + match client.response_get("usage").await { Ok(resp) => resp .json::() .await @@ -77,7 +77,8 @@ impl NanoGptProvider { NANOGPT_API_HOST.to_string() }; - let api_client = Self::build_client(&host, &api_key, tls_config)?; + let api_client = Self::build_client(&host, &api_key, tls_config)? + .with_request_builder(crate::session_context::session_id_request_builder()); Ok(Self { api_client, @@ -120,7 +121,7 @@ impl Provider for NanoGptProvider { async fn fetch_supported_models(&self) -> Result, ProviderError> { let response = self .api_client - .request(None, "models?detailed=true") + .request("models?detailed=true") .response_get() .await .map_err(|e| { @@ -176,7 +177,6 @@ impl Provider for NanoGptProvider { async fn stream( &self, model_config: &ModelConfig, - session_id: &str, system: &str, messages: &[Message], tools: &[Tool], @@ -196,7 +196,7 @@ impl Provider for NanoGptProvider { .with_retry(|| async { let resp = self .api_client - .response_post(Some(session_id), "chat/completions", &payload) + .response_post("chat/completions", &payload) .await?; handle_status(resp).await }) diff --git a/crates/goose/src/providers/ollama.rs b/crates/goose/src/providers/ollama.rs index 834e5ce72f30..e6a1ba71c98e 100644 --- a/crates/goose/src/providers/ollama.rs +++ b/crates/goose/src/providers/ollama.rs @@ -164,7 +164,8 @@ impl OllamaProvider { AuthMethod::NoAuth, timeout, tls_config, - )?; + )? + .with_request_builder(crate::session_context::session_id_request_builder()); Ok(Self { api_client, @@ -205,7 +206,8 @@ impl OllamaProvider { AuthMethod::NoAuth, timeout, tls_config, - )?; + )? + .with_request_builder(crate::session_context::session_id_request_builder()); if let Some(headers) = &config.headers { let mut header_map = reqwest::header::HeaderMap::new(); @@ -292,7 +294,6 @@ impl Provider for OllamaProvider { async fn stream( &self, model_config: &ModelConfig, - session_id: &str, system: &str, messages: &[Message], tools: &[Tool], @@ -312,7 +313,7 @@ impl Provider for OllamaProvider { .with_retry(|| async { let resp = self .api_client - .response_post(Some(session_id), "v1/chat/completions", &payload) + .response_post("v1/chat/completions", &payload) .await?; handle_status(resp).await }) @@ -326,7 +327,7 @@ impl Provider for OllamaProvider { async fn fetch_supported_models(&self) -> Result, ProviderError> { let response = self .api_client - .request(None, "api/tags") + .request("api/tags") .response_get() .await .map_err(|e| ProviderError::RequestFailed(format!("Failed to fetch models: {}", e)))?; diff --git a/crates/goose/src/providers/openai_def.rs b/crates/goose/src/providers/openai_def.rs index bf1615fdc21b..208d2cc7570c 100644 --- a/crates/goose/src/providers/openai_def.rs +++ b/crates/goose/src/providers/openai_def.rs @@ -113,7 +113,8 @@ pub async fn from_env( auth, std::time::Duration::from_secs(timeout_secs), tls_config, - )?; + )? + .with_request_builder(crate::session_context::session_id_request_builder()); if !parsed.query_params.is_empty() { api_client = api_client.with_query(parsed.query_params); @@ -258,7 +259,8 @@ pub fn from_custom_config( auth, std::time::Duration::from_secs(timeout_secs), tls_config, - )?; + )? + .with_request_builder(crate::session_context::session_id_request_builder()); if let Some(headers) = &config.headers { let mut header_map = reqwest::header::HeaderMap::new(); diff --git a/crates/goose/src/providers/openrouter.rs b/crates/goose/src/providers/openrouter.rs index c3a6a481d433..77b6e5413ac6 100644 --- a/crates/goose/src/providers/openrouter.rs +++ b/crates/goose/src/providers/openrouter.rs @@ -57,6 +57,7 @@ impl OpenRouterProvider { let auth = AuthMethod::BearerToken(api_key); let api_client = ApiClient::new_with_tls(host, auth, tls_config)? + .with_request_builder(crate::session_context::session_id_request_builder()) .with_header("HTTP-Referer", "https://goose-docs.ai")? .with_header("X-Title", "goose")?; @@ -195,7 +196,7 @@ impl Provider for OpenRouterProvider { async fn fetch_supported_models(&self) -> Result, ProviderError> { let response = self .api_client - .request(None, "api/v1/models") + .request("api/v1/models") .response_get() .await .map_err(|e| { @@ -242,11 +243,11 @@ impl Provider for OpenRouterProvider { async fn stream( &self, model_config: &ModelConfig, - session_id: &str, system: &str, messages: &[Message], tools: &[Tool], ) -> Result { + let session_id = crate::session_context::current_session_id().unwrap_or_default(); let mut payload = create_request( model_config, system, @@ -282,7 +283,7 @@ impl Provider for OpenRouterProvider { .with_retry(|| async { let resp = self .api_client - .response_post(Some(session_id), "api/v1/chat/completions", &payload) + .response_post("api/v1/chat/completions", &payload) .await?; handle_status(resp).await }) diff --git a/crates/goose/src/providers/provider_test.rs b/crates/goose/src/providers/provider_test.rs index 5284f2a5d3e5..b7ee3cad8a2c 100644 --- a/crates/goose/src/providers/provider_test.rs +++ b/crates/goose/src/providers/provider_test.rs @@ -26,15 +26,16 @@ pub async fn test_provider_configuration( vec![] }; - let mut stream = provider - .stream( + let mut stream = crate::session_context::with_session_id( + Some("test-session-id".to_string()), + provider.stream( &model_config, - "test-session-id", "You are an AI agent called goose. You use tools of connected extensions to solve problems.", &messages, &tools.into_iter().collect::>(), - ) - .await?; + ), + ) + .await?; let first_chunk = stream .next() diff --git a/crates/goose/src/providers/sagemaker_tgi.rs b/crates/goose/src/providers/sagemaker_tgi.rs index 18108bc26e25..e54f670379ad 100644 --- a/crates/goose/src/providers/sagemaker_tgi.rs +++ b/crates/goose/src/providers/sagemaker_tgi.rs @@ -318,15 +318,15 @@ impl Provider for SageMakerTgiProvider { async fn stream( &self, model_config: &ModelConfig, - session_id: &str, system: &str, messages: &[Message], tools: &[Tool], ) -> Result { + let session_id = crate::session_context::current_session_id().unwrap_or_default(); let session_id = if session_id.is_empty() { None } else { - Some(session_id) + Some(session_id.as_str()) }; let model_name = &model_config.model_name; diff --git a/crates/goose/src/providers/snowflake.rs b/crates/goose/src/providers/snowflake.rs index 0bb5c174043b..41db6df2b7f7 100644 --- a/crates/goose/src/providers/snowflake.rs +++ b/crates/goose/src/providers/snowflake.rs @@ -105,6 +105,7 @@ impl SnowflakeProvider { let auth = AuthMethod::BearerToken(token?); let api_client = ApiClient::new_with_tls(base_url, auth, tls_config)? + .with_request_builder(crate::session_context::session_id_request_builder()) .with_header("User-Agent", "goose")?; Ok(Self { @@ -114,14 +115,10 @@ impl SnowflakeProvider { }) } - async fn post( - &self, - session_id: Option<&str>, - payload: &Value, - ) -> Result { + async fn post(&self, payload: &Value) -> Result { let response = self .api_client - .response_post(session_id, "api/v2/cortex/inference:complete", payload) + .response_post("api/v2/cortex/inference:complete", payload) .await?; let status = response.status(); @@ -344,16 +341,10 @@ impl Provider for SnowflakeProvider { async fn stream( &self, model_config: &ModelConfig, - session_id: &str, system: &str, messages: &[Message], tools: &[Tool], ) -> Result { - let session_id = if session_id.is_empty() { - None - } else { - Some(session_id) - }; let payload = create_request(model_config, system, messages, tools)?; let mut log = start_log(model_config, &payload)?; @@ -361,7 +352,7 @@ impl Provider for SnowflakeProvider { let response = self .with_retry(|| async { let payload_clone = payload.clone(); - self.post(session_id, &payload_clone).await + self.post(&payload_clone).await }) .await?; diff --git a/crates/goose/src/providers/testprovider.rs b/crates/goose/src/providers/testprovider.rs index 6a6620edbc52..d6c2fd74cfd9 100644 --- a/crates/goose/src/providers/testprovider.rs +++ b/crates/goose/src/providers/testprovider.rs @@ -172,7 +172,6 @@ impl Provider for TestProvider { async fn stream( &self, model_config: &ModelConfig, - session_id: &str, system: &str, messages: &[Message], tools: &[Tool], @@ -181,9 +180,7 @@ impl Provider for TestProvider { if let Some(inner) = &self.inner { // Call inner provider's stream and collect it - let stream = inner - .stream(model_config, session_id, system, messages, tools) - .await?; + let stream = inner.stream(model_config, system, messages, tools).await?; let (message, usage) = super::base::collect_stream(stream).await?; let record = TestRecord { @@ -243,7 +240,6 @@ mod tests { async fn stream( &self, _model_config: &ModelConfig, - _session_id: &str, _system: &str, _messages: &[Message], _tools: &[Tool], @@ -281,13 +277,7 @@ mod tests { let model_config = ModelConfig::new("test-model"); let result = test_provider - .complete( - &model_config, - "test-session-id", - "You are helpful", - &[], - &[], - ) + .complete(&model_config, "You are helpful", &[], &[]) .await; assert!(result.is_ok()); @@ -306,13 +296,7 @@ mod tests { let model_config = ModelConfig::new("test-model"); let result = replay_provider - .complete( - &model_config, - "test-session-id", - "You are helpful", - &[], - &[], - ) + .complete(&model_config, "You are helpful", &[], &[]) .await; assert!(result.is_ok()); @@ -338,13 +322,7 @@ mod tests { let model_config = ModelConfig::new("test-model"); let result = replay_provider - .complete( - &model_config, - "test-session-id", - "Different system prompt", - &[], - &[], - ) + .complete(&model_config, "Different system prompt", &[], &[]) .await; assert!(result.is_err()); diff --git a/crates/goose/src/providers/tetrate.rs b/crates/goose/src/providers/tetrate.rs index 38ac05805069..3dddae36049f 100644 --- a/crates/goose/src/providers/tetrate.rs +++ b/crates/goose/src/providers/tetrate.rs @@ -57,6 +57,7 @@ impl TetrateProvider { let auth = AuthMethod::BearerToken(api_key); let api_client = ApiClient::new_with_tls(host, auth, tls_config)? + .with_request_builder(crate::session_context::session_id_request_builder()) .with_header("HTTP-Referer", "https://goose-docs.ai")? .with_header("X-Title", "goose")?; @@ -132,7 +133,6 @@ impl Provider for TetrateProvider { async fn stream( &self, model_config: &ModelConfig, - session_id: &str, system: &str, messages: &[Message], tools: &[Tool], @@ -152,7 +152,7 @@ impl Provider for TetrateProvider { .with_retry(|| async { let resp = self .api_client - .response_post(Some(session_id), "v1/chat/completions", &payload) + .response_post("v1/chat/completions", &payload) .await?; let resp = handle_status(resp) .await @@ -198,7 +198,7 @@ impl Provider for TetrateProvider { async fn fetch_supported_models(&self) -> Result, ProviderError> { let response = self .api_client - .response_get(None, "v1/models") + .response_get("v1/models") .await .map_err(|e| ProviderError::RequestFailed(e.to_string()))?; let json = handle_response_openai_compat(response).await?; diff --git a/crates/goose/src/providers/toolshim.rs b/crates/goose/src/providers/toolshim.rs index d16c51c53ff9..1707d629fdc4 100644 --- a/crates/goose/src/providers/toolshim.rs +++ b/crates/goose/src/providers/toolshim.rs @@ -581,7 +581,7 @@ impl LocalInterpreter { let request_messages = vec![Message::user().with_text(format_instruction)]; let mut stream = provider - .stream(&model_config, "toolshim-local", "", &request_messages, &[]) + .stream(&model_config, "", &request_messages, &[]) .await?; let mut content = String::new(); diff --git a/crates/goose/src/providers/xai.rs b/crates/goose/src/providers/xai.rs index 7488e713d353..71cbce32598f 100644 --- a/crates/goose/src/providers/xai.rs +++ b/crates/goose/src/providers/xai.rs @@ -64,7 +64,8 @@ impl ProviderDef for XaiProvider { .unwrap_or_else(|_| XAI_API_HOST.to_string()); let api_client = - ApiClient::new_with_tls(host, AuthMethod::BearerToken(api_key), tls_config)?; + ApiClient::new_with_tls(host, AuthMethod::BearerToken(api_key), tls_config)? + .with_request_builder(crate::session_context::session_id_request_builder()); Ok(OpenAiCompatibleProvider::new( XAI_PROVIDER_NAME.to_string(), diff --git a/crates/goose/src/providers/xai_oauth.rs b/crates/goose/src/providers/xai_oauth.rs index 8b675bd0e3ac..2cd66cd3ae80 100644 --- a/crates/goose/src/providers/xai_oauth.rs +++ b/crates/goose/src/providers/xai_oauth.rs @@ -706,13 +706,12 @@ impl Provider for XaiOAuthProvider { async fn stream( &self, model_config: &ModelConfig, - session_id: &str, system: &str, messages: &[Message], tools: &[Tool], ) -> Result { self.inner - .stream(model_config, session_id, system, messages, tools) + .stream(model_config, system, messages, tools) .await } @@ -792,7 +791,8 @@ impl ProviderDef for XaiOAuthProvider { host, AuthMethod::Custom(Box::new(SharedAuthProvider(auth_for_client))), tls_config, - )?; + )? + .with_request_builder(crate::session_context::session_id_request_builder()); let inner = OpenAiCompatibleProvider::new( XAI_OAUTH_PROVIDER_NAME.to_string(), diff --git a/crates/goose/src/security/adversary_inspector.rs b/crates/goose/src/security/adversary_inspector.rs index d0c038ca590e..a6d00f5790b4 100644 --- a/crates/goose/src/security/adversary_inspector.rs +++ b/crates/goose/src/security/adversary_inspector.rs @@ -326,16 +326,12 @@ impl AdversaryInspector { let model_config = resolve_model_config(&self.session_manager, session_id) .await .map_err(|e| anyhow::anyhow!("Could not resolve model config: {}", e))?; - let (response, _usage) = provider - .complete( - &model_config, - session_id, - system_prompt, - conversation.messages(), - &[], - ) - .await - .map_err(|e| anyhow::anyhow!("Adversary LLM call failed: {}", e))?; + let (response, _usage) = crate::session_context::with_session_id( + Some(session_id.to_string()), + provider.complete(&model_config, system_prompt, conversation.messages(), &[]), + ) + .await + .map_err(|e| anyhow::anyhow!("Adversary LLM call failed: {}", e))?; let output: String = response .content diff --git a/crates/goose/src/session/session_manager.rs b/crates/goose/src/session/session_manager.rs index ed64b8d1eb53..d9989bbab18f 100644 --- a/crates/goose/src/session/session_manager.rs +++ b/crates/goose/src/session/session_manager.rs @@ -2209,7 +2209,6 @@ mod tests { async fn stream( &self, _model_config: &ModelConfig, - _session_id: &str, _system: &str, _messages: &[Message], _tools: &[rmcp::model::Tool], @@ -2220,7 +2219,6 @@ mod tests { async fn complete( &self, _model_config: &ModelConfig, - _session_id: &str, _system: &str, _messages: &[Message], _tools: &[Tool], diff --git a/crates/goose/src/session_context.rs b/crates/goose/src/session_context.rs index ad4046f692e5..4c391485f7ae 100644 --- a/crates/goose/src/session_context.rs +++ b/crates/goose/src/session_context.rs @@ -1,10 +1,11 @@ -use tokio::task_local; +use reqwest::header::{HeaderName, HeaderValue}; pub const SESSION_ID_HEADER: &str = "agent-session-id"; + pub const TOOL_CALL_REQUEST_ID_HEADER: &str = "agent-tool-call-request-id"; pub const WORKING_DIR_HEADER: &str = "agent-working-dir"; -task_local! { +tokio::task_local! { pub static SESSION_ID: Option; } @@ -12,17 +13,29 @@ pub async fn with_session_id(session_id: Option, f: F) -> F::Output where F: std::future::Future, { - if let Some(id) = session_id { - SESSION_ID.scope(Some(id), f).await - } else { - f.await - } + SESSION_ID.scope(session_id, f).await } pub fn current_session_id() -> Option { SESSION_ID.try_with(|id| id.clone()).ok().flatten() } +pub fn session_id_request_builder() -> goose_providers::api_client::RequestBuilderDecorator { + std::sync::Arc::new(|request| { + let (client, request) = request.build_split(); + let mut request = request?; + let session_header = HeaderName::from_static(SESSION_ID_HEADER); + request.headers_mut().remove(&session_header); + + if let Some(session_id) = current_session_id() { + let value = HeaderValue::from_str(&session_id)?; + request.headers_mut().insert(session_header, value); + } + + Ok(reqwest::RequestBuilder::from_parts(client, request)) + }) +} + /// Local OS user running goose, shared by the OTLP `user.name` resource /// attribute and the `session.user` span attribute so the two never drift. pub fn session_user() -> String { @@ -63,6 +76,21 @@ mod tests { .await; } + #[tokio::test] + async fn test_session_id_none_clears_outer_scope() { + with_session_id(Some("outer-session".to_string()), async { + assert_eq!(current_session_id(), Some("outer-session".to_string())); + + with_session_id(None, async { + assert_eq!(current_session_id(), None); + }) + .await; + + assert_eq!(current_session_id(), Some("outer-session".to_string())); + }) + .await; + } + #[tokio::test] async fn test_session_id_scoped_correctly() { assert_eq!(current_session_id(), None); diff --git a/crates/goose/tests/acp_custom_requests_test.rs b/crates/goose/tests/acp_custom_requests_test.rs index d88356533465..36c13e8aa767 100644 --- a/crates/goose/tests/acp_custom_requests_test.rs +++ b/crates/goose/tests/acp_custom_requests_test.rs @@ -60,7 +60,6 @@ impl Provider for MockProvider { async fn stream( &self, _model_config: &ModelConfig, - _session_id: &str, _system: &str, _messages: &[goose::conversation::message::Message], _tools: &[rmcp::model::Tool], diff --git a/crates/goose/tests/acp_fixtures/provider.rs b/crates/goose/tests/acp_fixtures/provider.rs index d5c62f3e6c06..595d9f0028c5 100644 --- a/crates/goose/tests/acp_fixtures/provider.rs +++ b/crates/goose/tests/acp_fixtures/provider.rs @@ -75,9 +75,11 @@ impl AcpProviderSession { .get(session_id.as_ref()) .cloned() .unwrap_or_else(|| ModelConfig::new(TEST_MODEL)); - let mut stream = provider - .stream(&model_config, &session_id, "", &[message], &[]) - .await?; + let mut stream = goose::session_context::with_session_id( + Some(session_id.to_string()), + provider.stream(&model_config, "", &[message], &[]), + ) + .await?; let mut text = String::new(); let mut tool_error = false; let mut saw_tool = false; diff --git a/crates/goose/tests/acp_secret_cache_invalidation_test.rs b/crates/goose/tests/acp_secret_cache_invalidation_test.rs index 9f4bf88f9eb1..ee4100ece9e8 100644 --- a/crates/goose/tests/acp_secret_cache_invalidation_test.rs +++ b/crates/goose/tests/acp_secret_cache_invalidation_test.rs @@ -28,7 +28,6 @@ impl Provider for MockProvider { async fn stream( &self, _model_config: &ModelConfig, - _session_id: &str, _system: &str, _messages: &[goose::conversation::message::Message], _tools: &[rmcp::model::Tool], diff --git a/crates/goose/tests/agent.rs b/crates/goose/tests/agent.rs index 140d2b981eb4..c43cc9875ca8 100644 --- a/crates/goose/tests/agent.rs +++ b/crates/goose/tests/agent.rs @@ -526,7 +526,6 @@ mod tests { async fn stream( &self, _model_config: &ModelConfig, - _session_id: &str, _system_prompt: &str, _messages: &[Message], _tools: &[Tool], @@ -698,7 +697,6 @@ mod tests { async fn stream( &self, _model_config: &ModelConfig, - _session_id: &str, _system_prompt: &str, _messages: &[Message], _tools: &[Tool], @@ -879,7 +877,6 @@ mod tests { async fn stream( &self, _model_config: &ModelConfig, - _session_id: &str, system_prompt: &str, _messages: &[Message], _tools: &[Tool], @@ -1232,7 +1229,6 @@ mod tests { async fn stream( &self, _model_config: &ModelConfig, - _session_id: &str, _system_prompt: &str, _messages: &[Message], _tools: &[Tool], @@ -1504,7 +1500,6 @@ mod tests { async fn stream( &self, _model_config: &ModelConfig, - _session_id: &str, _system_prompt: &str, _messages: &[Message], _tools: &[Tool], @@ -1705,7 +1700,6 @@ mod tests { async fn stream( &self, _model_config: &ModelConfig, - _session_id: &str, _system_prompt: &str, _messages: &[Message], _tools: &[Tool], @@ -1857,7 +1851,6 @@ mod tests { async fn stream( &self, _model_config: &ModelConfig, - _session_id: &str, _system_prompt: &str, _messages: &[Message], _tools: &[Tool], @@ -2063,7 +2056,6 @@ mod tests { async fn stream( &self, _model_config: &ModelConfig, - _session_id: &str, _system_prompt: &str, _messages: &[Message], _tools: &[Tool], @@ -2415,7 +2407,6 @@ mod tests { async fn stream( &self, _model_config: &ModelConfig, - _session_id: &str, _system_prompt: &str, _messages: &[Message], _tools: &[Tool], diff --git a/crates/goose/tests/compaction.rs b/crates/goose/tests/compaction.rs index 5d4d0769f72e..f38a1966ecaa 100644 --- a/crates/goose/tests/compaction.rs +++ b/crates/goose/tests/compaction.rs @@ -101,7 +101,6 @@ impl Provider for MockCompactionProvider { async fn stream( &self, _model_config: &ModelConfig, - _session_id: &str, system_prompt: &str, messages: &[Message], _tools: &[Tool], diff --git a/crates/goose/tests/local_inference_integration.rs b/crates/goose/tests/local_inference_integration.rs index 431aeb277438..174b871fbebd 100644 --- a/crates/goose/tests/local_inference_integration.rs +++ b/crates/goose/tests/local_inference_integration.rs @@ -37,7 +37,7 @@ async fn test_local_inference_stream_produces_output() { let messages = vec![Message::user().with_text("Say hello.")]; let mut stream = provider - .stream(&model_config, "test-session", system, &messages, &[]) + .stream(&model_config, system, &messages, &[]) .await .expect("stream should start"); @@ -82,7 +82,7 @@ async fn test_local_inference_large_prompt() { let start = std::time::Instant::now(); let (response, _usage) = provider - .complete(&model_config, "test-session", "", &messages, &[]) + .complete(&model_config, "", &messages, &[]) .await .expect("large prompt completion should succeed"); let elapsed = start.elapsed(); @@ -149,7 +149,7 @@ async fn test_local_inference_vision_produces_output() { .with_image(image_b64, "image/png")]; let mut stream = provider - .stream(&model_config, "test-vision-session", system, &messages, &[]) + .stream(&model_config, system, &messages, &[]) .await .expect("stream should start for vision input"); @@ -194,7 +194,7 @@ async fn test_local_inference_vision_text_only_model_graceful() { .with_image(image_b64, "image/png")]; let mut stream = provider - .stream(&model_config, "test-session", system, &messages, &[]) + .stream(&model_config, system, &messages, &[]) .await .expect("stream should start"); diff --git a/crates/goose/tests/local_inference_perf.rs b/crates/goose/tests/local_inference_perf.rs index bb9ba4fbe65f..44250ab97dc9 100644 --- a/crates/goose/tests/local_inference_perf.rs +++ b/crates/goose/tests/local_inference_perf.rs @@ -33,7 +33,7 @@ async fn test_local_inference_cold_vs_warm() { let messages = vec![Message::user().with_text("What is 2+2?")]; let start = Instant::now(); let (response, _) = provider - .complete(&model_config, "perf-session", "", &messages, &[]) + .complete(&model_config, "", &messages, &[]) .await .expect("cold completion should succeed"); let cold_elapsed = start.elapsed(); @@ -46,7 +46,7 @@ async fn test_local_inference_cold_vs_warm() { let messages2 = vec![Message::user().with_text("What is 3+3?")]; let start2 = Instant::now(); let (response2, _) = provider - .complete(&model_config, "perf-session", "", &messages2, &[]) + .complete(&model_config, "", &messages2, &[]) .await .expect("warm completion should succeed"); let warm_elapsed = start2.elapsed(); diff --git a/crates/goose/tests/mcp_integration_test.rs b/crates/goose/tests/mcp_integration_test.rs index fc3ed1fd4ff2..1c7be139bfe7 100644 --- a/crates/goose/tests/mcp_integration_test.rs +++ b/crates/goose/tests/mcp_integration_test.rs @@ -75,7 +75,6 @@ impl Provider for MockProvider { async fn stream( &self, _model_config: &ModelConfig, - _session_id: &str, _system: &str, _messages: &[Message], _tools: &[Tool], diff --git a/crates/goose/tests/providers.rs b/crates/goose/tests/providers.rs index 0231e60b1d15..9b1411563d8e 100644 --- a/crates/goose/tests/providers.rs +++ b/crates/goose/tests/providers.rs @@ -287,7 +287,7 @@ impl ProviderFixture { provider, model_config, agent, - session_id, + session_id: session_id.to_string(), _mcp: mcp, _guard: guard, _temp_dir: temp_dir, @@ -318,16 +318,16 @@ impl ProviderFixture { let message = Message::user().with_text(prompt); let model_config = model_config.unwrap_or_else(|| self.model_config.clone()); - let (response1, _) = self - .provider - .complete( + let (response1, _) = goose::session_context::with_session_id( + Some(self.session_id.clone()), + self.provider.complete( &model_config, - &self.session_id, &system, std::slice::from_ref(&message), &tools, - ) - .await?; + ), + ) + .await?; // Agentic CLI providers (claude-code, codex) call tools internally and // return the final text result directly — no tool_request in the response. @@ -364,16 +364,16 @@ impl ProviderFixture { .unwrap(); let tool_response = Message::user().with_tool_response(&tool_req.id, Ok(result)); - let (response2, _) = self - .provider - .complete( + let (response2, _) = goose::session_context::with_session_id( + Some(self.session_id.clone()), + self.provider.complete( &model_config, - &self.session_id, &system, &[message, response1, tool_response], &tools, - ) - .await?; + ), + ) + .await?; Ok(response2) } @@ -381,16 +381,16 @@ impl ProviderFixture { let message = Message::user().with_text("Just say hello!"); let model_config = self.model_config.clone(); - let (response, _) = self - .provider - .complete( + let (response, _) = goose::session_context::with_session_id( + Some(self.session_id.clone()), + self.provider.complete( &model_config, - &self.session_id, "You are a helpful assistant.", &[message], &[], - ) - .await?; + ), + ) + .await?; assert!(!response.content.is_empty()); assert!(response @@ -422,16 +422,16 @@ impl ProviderFixture { let messages = vec![Message::user().with_text(&large_message_content)]; let model_config = self.model_config.clone(); - let result = self - .provider - .complete( + let result = goose::session_context::with_session_id( + Some(self.session_id.clone()), + self.provider.complete( &model_config, - &self.session_id, "You are a helpful assistant.", &messages, &[], - ) - .await; + ), + ) + .await; println!("=== {}::context_length_exceeded_error ===", self.name); dbg!(&result); @@ -476,16 +476,12 @@ impl ProviderFixture { goose_providers::model::ModelConfig::new(alt).with_canonical_limits(&self.name); let message = Message::user().with_text("Just say hello!"); - let (response, _) = self - .provider - .complete( - &alt_config, - &self.session_id, - "You are a helpful assistant.", - &[message], - &[], - ) - .await?; + let (response, _) = goose::session_context::with_session_id( + Some(self.session_id.clone()), + self.provider + .complete(&alt_config, "You are a helpful assistant.", &[message], &[]), + ) + .await?; assert!(response .content diff --git a/crates/goose/tests/session_id_propagation_test.rs b/crates/goose/tests/session_id_propagation_test.rs index dae559f1326a..0666fe448905 100644 --- a/crates/goose/tests/session_id_propagation_test.rs +++ b/crates/goose/tests/session_id_propagation_test.rs @@ -2,7 +2,7 @@ use goose::conversation::message::Message; use goose::providers::api_client::{ApiClient, AuthMethod}; use goose::providers::base::Provider; use goose::providers::openai::OpenAiProvider; -use goose::session_context::SESSION_ID_HEADER; +use goose::session_context::{session_id_request_builder, SESSION_ID_HEADER}; use goose_providers::model::ModelConfig; use serde_json::json; use std::sync::Arc; @@ -41,7 +41,8 @@ fn create_test_provider(mock_server_url: &str) -> Box { AuthMethod::BearerToken("test-key".to_string()), None, ) - .unwrap(); + .unwrap() + .with_request_builder(session_id_request_builder()); Box::new(OpenAiProvider::new(api_client)) } @@ -144,16 +145,17 @@ async fn setup_mock_server() -> (MockServer, HeaderCapture, Box) { async fn make_request(provider: &dyn Provider, session_id: &str) { let message = Message::user().with_text("test message"); let model_config = ModelConfig::new("gpt-5-nano"); - let _ = provider - .complete( + let _ = goose::session_context::with_session_id( + Some(session_id.to_string()), + provider.complete( &model_config, - session_id, "You are a helpful assistant.", &[message], &[], - ) - .await - .unwrap(); + ), + ) + .await + .unwrap(); } #[tokio::test] diff --git a/crates/goose/tests/tetrate_streaming.rs b/crates/goose/tests/tetrate_streaming.rs index a065a10083ef..fded3613d5a0 100644 --- a/crates/goose/tests/tetrate_streaming.rs +++ b/crates/goose/tests/tetrate_streaming.rs @@ -30,7 +30,6 @@ mod tetrate_streaming_tests { let mut stream = provider .stream( &model_config, - "test-session-id", "You are a helpful assistant that counts numbers.", &messages, &[], @@ -105,7 +104,6 @@ mod tetrate_streaming_tests { let mut stream = provider .stream( &model_config, - "test-session-id", "You are a helpful assistant with access to weather information.", &messages, &[weather_tool], @@ -156,7 +154,6 @@ mod tetrate_streaming_tests { let mut stream = provider .stream( &model_config, - "test-session-id", "You are a helpful assistant.", &messages, &[], @@ -194,7 +191,6 @@ mod tetrate_streaming_tests { let mut stream = provider .stream( &model_config, - "test-session-id", "You are a helpful assistant that writes detailed essays.", &messages, &[], @@ -256,7 +252,6 @@ mod tetrate_streaming_tests { let result = provider .stream( &model_config, - "test-session-id", "You are a helpful assistant.", &messages, &[], @@ -287,7 +282,6 @@ mod tetrate_streaming_tests { let stream1 = provider .stream( &model_config, - "test-session-id", "You are a helpful assistant.", &messages1, &[], @@ -297,7 +291,6 @@ mod tetrate_streaming_tests { let stream2 = provider .stream( &model_config, - "test-session-id", "You are a helpful assistant.", &messages2, &[],