Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
4 changes: 4 additions & 0 deletions crates/forge_repo/src/provider/anthropic.rs
Original file line number Diff line number Diff line change
Expand Up @@ -87,6 +87,10 @@ impl<H: HttpInfra> Anthropic<H> {
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
}
}
Expand Down
22 changes: 19 additions & 3 deletions crates/forge_repo/src/provider/google.rs
Original file line number Diff line number Diff line change
Expand Up @@ -20,6 +20,7 @@ struct Google<T> {
chat_url: Url,
models: forge_domain::ModelSource<Url>,
use_api_key_header: bool,
custom_headers: Vec<(String, String)>,
}

impl<H: HttpInfra> Google<H> {
Expand All @@ -30,7 +31,14 @@ impl<H: HttpInfra> Google<H> {
models: forge_domain::ModelSource<Url>,
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)> {
Expand All @@ -45,6 +53,7 @@ impl<H: HttpInfra> Google<H> {
));
}

headers.extend(self.custom_headers.iter().cloned());
headers
}
}
Expand Down Expand Up @@ -179,13 +188,20 @@ impl<F: HttpInfra> GoogleResponseRepository<F> {
}
};

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]
Expand Down
5 changes: 5 additions & 0 deletions crates/forge_repo/src/provider/openai_responses/repository.rs
Original file line number Diff line number Diff line change
Expand Up @@ -45,6 +45,7 @@ impl<H: HttpInfra> OpenAIResponsesProvider<H> {

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,
Expand Down Expand Up @@ -123,6 +124,10 @@ impl<H: HttpInfra> OpenAIResponsesProvider<H> {
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.
//
Expand Down
257 changes: 254 additions & 3 deletions crates/forge_repo/src/provider/opencode.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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;
Expand Down Expand Up @@ -61,9 +61,23 @@ impl<F: HttpInfra + EnvironmentInfra<Config = forge_config::ForgeConfig> + 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<Url>, model_id: &ModelId) -> Provider<Url> {
fn build_provider(
&self,
provider: &Provider<Url>,
model_id: &ModelId,
conversation_id: ConversationId,
) -> Provider<Url> {
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 {
Expand Down Expand Up @@ -99,7 +113,12 @@ impl<F: HttpInfra + EnvironmentInfra<Config = forge_config::ForgeConfig> + Sync>
provider: Provider<Url>,
) -> ResultStream<ChatCompletionMessage, anyhow::Error> {
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 => {
Expand Down Expand Up @@ -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<Vec<(Url, HeaderMap, Bytes)>>,
}

impl EnvironmentInfra for RecordingInfra {
type Config = forge_config::ForgeConfig;

fn get_config(&self) -> anyhow::Result<Self::Config> {
Ok(Self::Config::default())
}

fn get_environment(&self) -> forge_domain::Environment {
use fake::{Fake, Faker};
Faker.fake()
}

fn get_env_var(&self, _: &str) -> Option<String> {
None
}

fn get_env_vars(&self) -> BTreeMap<String, String> {
BTreeMap::new()
}

async fn update_environment(&self, _: Vec<forge_domain::ConfigOperation>) -> Result<()> {
Ok(())
}
}

#[async_trait::async_trait]
impl HttpInfra for RecordingInfra {
async fn http_delete(&self, _: &Url) -> Result<reqwest::Response> {
anyhow::bail!("Unexpected DELETE")
}

async fn http_get(&self, _: &Url, _: Option<HeaderMap>) -> Result<reqwest::Response> {
anyhow::bail!("Unexpected GET")
}

async fn http_post(
&self,
url: &Url,
headers: Option<HeaderMap>,
body: Bytes,
) -> Result<reqwest::Response> {
self.requests
.lock()
.unwrap()
.push((url.clone(), headers.unwrap(), body));
anyhow::bail!("Recorded request")
}

async fn http_eventsource(
&self,
url: &Url,
headers: Option<HeaderMap>,
body: Bytes,
) -> Result<EventSource> {
self.requests
.lock()
.unwrap()
.push((url.clone(), headers.unwrap(), body));
anyhow::bail!("Recorded request")
}
}

fn fixture_provider(id: ProviderId) -> anyhow::Result<Provider<Url>> {
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::<Vec<_>>();
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::<Vec<_>>();
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::<std::collections::HashSet<_>>()
.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::<Vec<_>>();
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-") {
Expand Down
Loading