diff --git a/crates/forge_repo/src/provider/anthropic.rs b/crates/forge_repo/src/provider/anthropic.rs index 72fb10d46c..896c378d49 100644 --- a/crates/forge_repo/src/provider/anthropic.rs +++ b/crates/forge_repo/src/provider/anthropic.rs @@ -87,6 +87,10 @@ impl Anthropic { headers.push(("anthropic-beta".to_string(), betas.join(","))); } + if let Some(custom_headers) = &self.provider.custom_headers { + headers.extend(custom_headers.iter().map(|(k, v)| (k.clone(), v.clone()))); + } + headers } } diff --git a/crates/forge_repo/src/provider/google.rs b/crates/forge_repo/src/provider/google.rs index 88e817116f..51aaf6275a 100644 --- a/crates/forge_repo/src/provider/google.rs +++ b/crates/forge_repo/src/provider/google.rs @@ -20,6 +20,7 @@ struct Google { chat_url: Url, models: forge_domain::ModelSource, use_api_key_header: bool, + custom_headers: Vec<(String, String)>, } impl Google { @@ -30,7 +31,14 @@ impl Google { models: forge_domain::ModelSource, use_api_key_header: bool, ) -> Self { - Self { http, api_key, chat_url, models, use_api_key_header } + Self { + http, + api_key, + chat_url, + models, + use_api_key_header, + custom_headers: Vec::new(), + } } fn get_headers(&self) -> Vec<(String, String)> { @@ -45,6 +53,7 @@ impl Google { )); } + headers.extend(self.custom_headers.iter().cloned()); headers } } @@ -179,13 +188,20 @@ impl GoogleResponseRepository { } }; - Ok(Google::new( + let mut client = Google::new( self.infra.clone(), token, chat_url, models, use_api_key_header, - )) + ); + if let Some(custom_headers) = &provider.custom_headers { + client.custom_headers = custom_headers + .iter() + .map(|(k, v)| (k.clone(), v.clone())) + .collect(); + } + Ok(client) } } #[async_trait::async_trait] diff --git a/crates/forge_repo/src/provider/openai_responses/repository.rs b/crates/forge_repo/src/provider/openai_responses/repository.rs index 2cd34cbed4..d09bab99a8 100644 --- a/crates/forge_repo/src/provider/openai_responses/repository.rs +++ b/crates/forge_repo/src/provider/openai_responses/repository.rs @@ -45,6 +45,7 @@ impl OpenAIResponsesProvider { if provider.id == ProviderId::CODEX || provider.id == ProviderId::OPENCODE_ZEN + || provider.id == ProviderId::OPENCODE_GO || provider.id == ProviderId::OPENAI_RESPONSES_COMPATIBLE { // These providers already configure a complete Responses endpoint, @@ -123,6 +124,10 @@ impl OpenAIResponsesProvider { forge_domain::AuthMethod::AwsProfile => {} }); + if let Some(custom_headers) = &self.provider.custom_headers { + headers.extend(custom_headers.iter().map(|(k, v)| (k.clone(), v.clone()))); + } + // Codex provider requires the ChatGPT-Account-Id header extracted // from the JWT at login. // diff --git a/crates/forge_repo/src/provider/opencode.rs b/crates/forge_repo/src/provider/opencode.rs index dc534f4062..03445e9700 100644 --- a/crates/forge_repo/src/provider/opencode.rs +++ b/crates/forge_repo/src/provider/opencode.rs @@ -6,7 +6,7 @@ use forge_app::domain::{ ResultStream, }; use forge_app::{EnvironmentInfra, HttpInfra}; -use forge_domain::ChatRepository; +use forge_domain::{ChatRepository, ConversationId}; use url::Url; use crate::provider::anthropic::AnthropicResponseRepository; @@ -61,9 +61,23 @@ impl + Sync> /// Derives the endpoint URL from the provider's configured base URL so that /// both OpenCode Zen and OpenCode Go (and any future variants) are routed /// to their correct endpoints. - fn build_provider(&self, provider: &Provider, model_id: &ModelId) -> Provider { + fn build_provider( + &self, + provider: &Provider, + model_id: &ModelId, + conversation_id: ConversationId, + ) -> Provider { let backend = self.get_backend(model_id); let mut new_provider = provider.clone(); + // This clone belongs to one chat request, never to the shared provider + // config. Keep session metadata outside the body + // transformations used by each adapter. + let headers = new_provider.custom_headers.get_or_insert_default(); + headers.retain(|name, _| !name.eq_ignore_ascii_case("x-opencode-session")); + headers.insert( + "x-opencode-session".to_string(), + conversation_id.to_string(), + ); let base = provider.url.as_str().trim_end_matches('/'); match backend { @@ -99,7 +113,12 @@ impl + Sync> provider: Provider, ) -> ResultStream { let backend = self.get_backend(model_id); - let adapted_provider = self.build_provider(&provider, model_id); + // Standalone requests without a conversation still need a session, but + // must not share a provider-wide ID with unrelated requests. + let conversation_id = context + .conversation_id + .unwrap_or_else(ConversationId::generate); + let adapted_provider = self.build_provider(&provider, model_id, conversation_id); match backend { OpenCodeBackend::Anthropic => { @@ -153,10 +172,242 @@ enum OpenCodeBackend { #[cfg(test)] mod tests { + use std::collections::{BTreeMap, HashMap}; + use std::sync::Mutex; + + use bytes::Bytes; + use forge_domain::{AuthCredential, AuthDetails, ContextMessage, ProviderId}; + use forge_eventsource::EventSource; use pretty_assertions::assert_eq; + use reqwest::header::HeaderMap; use super::*; + #[derive(Default)] + struct RecordingInfra { + requests: Mutex>, + } + + impl EnvironmentInfra for RecordingInfra { + type Config = forge_config::ForgeConfig; + + fn get_config(&self) -> anyhow::Result { + Ok(Self::Config::default()) + } + + fn get_environment(&self) -> forge_domain::Environment { + use fake::{Fake, Faker}; + Faker.fake() + } + + fn get_env_var(&self, _: &str) -> Option { + None + } + + fn get_env_vars(&self) -> BTreeMap { + BTreeMap::new() + } + + async fn update_environment(&self, _: Vec) -> Result<()> { + Ok(()) + } + } + + #[async_trait::async_trait] + impl HttpInfra for RecordingInfra { + async fn http_delete(&self, _: &Url) -> Result { + anyhow::bail!("Unexpected DELETE") + } + + async fn http_get(&self, _: &Url, _: Option) -> Result { + anyhow::bail!("Unexpected GET") + } + + async fn http_post( + &self, + url: &Url, + headers: Option, + body: Bytes, + ) -> Result { + self.requests + .lock() + .unwrap() + .push((url.clone(), headers.unwrap(), body)); + anyhow::bail!("Recorded request") + } + + async fn http_eventsource( + &self, + url: &Url, + headers: Option, + body: Bytes, + ) -> Result { + self.requests + .lock() + .unwrap() + .push((url.clone(), headers.unwrap(), body)); + anyhow::bail!("Recorded request") + } + } + + fn fixture_provider(id: ProviderId) -> anyhow::Result> { + Ok(Provider { + id: id.clone(), + provider_type: Default::default(), + response: Some(ProviderResponse::OpenCode), + url: Url::parse("https://opencode.ai/zen/go")?, + models: Some(forge_domain::ModelSource::Hardcoded(vec![])), + auth_methods: vec![forge_domain::AuthMethod::ApiKey], + url_params: vec![], + credential: Some(AuthCredential { + id, + auth_details: AuthDetails::ApiKey("fixture-key".to_string().into()), + url_params: HashMap::new(), + }), + custom_headers: Some(HashMap::from([ + ( + "X-OpenCode-Session".to_string(), + "stale-static-session".to_string(), + ), + ("x-custom".to_string(), "preserved".to_string()), + ])), + }) + } + + #[tokio::test] + async fn test_session_reaches_every_opencode_adapter() { + for provider_id in [ProviderId::OPENCODE_GO, ProviderId::OPENCODE_ZEN] { + let fixture = fixture_provider(provider_id).unwrap(); + let original = fixture.clone(); + let infra = Arc::new(RecordingInfra::default()); + let repo = OpenCodeZenResponseRepository::new(infra.clone()); + let conversation_id = ConversationId::generate(); + let other_id = ConversationId::generate(); + let context = ChatContext::default() + .conversation_id(conversation_id) + .add_message(ContextMessage::user("hello", None)); + let models = [ + ("deepseek-v4-flash", "/zen/go/v1/chat/completions"), + ("gpt-5", "/zen/go/v1/responses"), + ("claude-sonnet-4-5", "/zen/go/v1/messages"), + ( + "gemini-3-flash", + "/zen/go/v1/models/gemini-3-flash:streamGenerateContent", + ), + ]; + + // Repeated turns/retries and resumed contexts use the same ID; an + // interleaved conversation must not overwrite that session. + for id in [conversation_id, other_id, conversation_id] { + for (model, _) in models { + let result = repo + .chat( + &ModelId::new(model), + context.clone().conversation_id(id), + fixture.clone(), + ) + .await; + assert!(result.is_err()); + } + } + let requests = infra.requests.lock().unwrap(); + let actual = requests + .iter() + .map(|(url, headers, body)| { + let json: serde_json::Value = serde_json::from_slice(body).unwrap(); + assert!(json.get("session_id").is_none()); + ( + url.path().to_string(), + headers + .get("x-opencode-session") + .map(|value| value.to_str().unwrap().to_string()), + headers + .get("x-custom") + .map(|value| value.to_str().unwrap().to_string()), + headers.get_all("x-opencode-session").iter().count(), + ) + }) + .collect::>(); + let expected = [conversation_id, other_id, conversation_id] + .into_iter() + .flat_map(|id| { + models.map(|(_, path)| { + ( + path.to_string(), + Some(id.to_string()), + Some("preserved".to_string()), + 1, + ) + }) + }) + .collect::>(); + assert_eq!(actual, expected); + assert_eq!(fixture, original); + } + } + + #[tokio::test] + async fn test_standalone_opencode_requests_have_independent_sessions() { + let fixture = fixture_provider(ProviderId::OPENCODE_GO).unwrap(); + let infra = Arc::new(RecordingInfra::default()); + let repo = OpenCodeZenResponseRepository::new(infra.clone()); + for _ in 0..2 { + let result = repo + .chat( + &ModelId::new("deepseek-v4-flash"), + ChatContext::default(), + fixture.clone(), + ) + .await; + assert!(result.is_err()); + } + let requests = infra.requests.lock().unwrap(); + let actual = requests + .iter() + .map(|(_, headers, _)| { + ConversationId::parse(headers["x-opencode-session"].to_str().unwrap()).unwrap() + }) + .collect::>() + .len(); + let expected = 2; + assert_eq!(actual, expected); + } + + #[tokio::test] + async fn test_other_providers_do_not_receive_opencode_session() { + let mut fixture = fixture_provider(ProviderId::OPENAI).unwrap(); + fixture.custom_headers = None; + let infra = Arc::new(RecordingInfra::default()); + let context = ChatContext::default().conversation_id(ConversationId::generate()); + let model = ModelId::new("fixture-model"); + + let result = OpenAIResponseRepository::new(infra.clone()) + .chat(&model, context.clone(), fixture.clone()) + .await; + assert!(result.is_err()); + let result = OpenAIResponsesResponseRepository::new(infra.clone()) + .chat(&model, context.clone(), fixture.clone()) + .await; + assert!(result.is_err()); + let result = AnthropicResponseRepository::new(infra.clone()) + .chat(&model, context.clone(), fixture.clone()) + .await; + assert!(result.is_err()); + let result = GoogleResponseRepository::new(infra.clone()) + .chat(&model, context, fixture) + .await; + assert!(result.is_err()); + let actual = infra + .requests + .lock() + .unwrap() + .iter() + .map(|(_, headers, _)| headers.contains_key("x-opencode-session")) + .collect::>(); + let expected = vec![false, false, false, false]; + assert_eq!(actual, expected); + } + /// Helper function to determine backend routing (mirrors get_backend logic) fn get_backend_for_test(model_id: &str) -> OpenCodeBackend { if model_id.starts_with("claude-") {