Skip to content
Merged
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
57 changes: 0 additions & 57 deletions crates/agent/src/tests/mod.rs
Original file line number Diff line number Diff line change
Expand Up @@ -4069,63 +4069,6 @@ async fn test_send_retry_on_error(cx: &mut TestAppContext) {
});
}

#[gpui::test]
async fn test_send_retry_on_http_send_error(cx: &mut TestAppContext) {
let ThreadTest { thread, model, .. } = setup(cx, TestModel::Fake).await;
let fake_model = model.as_fake();

let mut events = thread
.update(cx, |thread, cx| {
thread.send(UserMessageId::new(), ["Hello!"], cx)
})
.expect("thread send should start");
cx.run_until_parked();

fake_model.send_last_completion_stream_error(LanguageModelCompletionError::HttpSend {
provider: LanguageModelProviderName::new("OpenAI"),
error: anyhow::anyhow!("response headers timed out after 10s"),
});
fake_model.end_last_completion_stream();

cx.executor().advance_clock(BASE_RETRY_DELAY);
cx.run_until_parked();

fake_model.send_last_completion_stream_text_chunk("Recovered!");
fake_model.end_last_completion_stream();
cx.run_until_parked();

let mut retry_events = Vec::new();
while let Some(Ok(event)) = events.next().await {
match event {
ThreadEvent::Retry(retry_status) => {
retry_events.push(retry_status);
}
ThreadEvent::Stop(..) => break,
_ => {}
}
}

assert_eq!(retry_events.len(), 1);
assert!(matches!(
retry_events[0],
acp_thread::RetryStatus { attempt: 1, .. }
));
thread.read_with(cx, |thread, _cx| {
assert_eq!(
thread.to_markdown(),
indoc! {"
## User

Hello!

## Assistant

Recovered!
"}
)
});
}

#[gpui::test]
async fn test_send_retry_finishes_tool_calls_on_error(cx: &mut TestAppContext) {
let ThreadTest { thread, model, .. } = setup(cx, TestModel::Fake).await;
Expand Down
204 changes: 12 additions & 192 deletions crates/language_models/src/provider/openai_subscribed.rs
Original file line number Diff line number Diff line change
Expand Up @@ -12,15 +12,12 @@ use language_model::{
LanguageModelProviderName, LanguageModelProviderState, LanguageModelRequest,
LanguageModelToolChoice, RateLimiter,
};
use open_ai::{
ReasoningEffort,
responses::{StreamResponseOptions, stream_response_with_options},
};
use open_ai::{ReasoningEffort, responses::stream_response};
use rand::RngCore as _;
use serde::{Deserialize, Serialize};
use sha2::{Digest, Sha256};
use std::sync::Arc;
use std::time::{Duration, SystemTime, UNIX_EPOCH};
use std::time::{SystemTime, UNIX_EPOCH};
use ui::{ConfiguredApiCard, prelude::*};
use url::form_urlencoded;
use util::ResultExt as _;
Expand All @@ -38,31 +35,6 @@ const CLIENT_ID: &str = "app_EMoamEEZ73f0CkXaXp7hrann";

const CREDENTIALS_KEY: &str = "https://chatgpt.com/backend-api/codex";
const TOKEN_REFRESH_BUFFER_MS: u64 = 5 * 60 * 1000;
const CODEX_RESPONSE_HEADER_TIMEOUT: Duration = Duration::from_secs(10);

fn codex_extra_headers(
account_id: Option<&str>,
session_id: Option<&str>,
) -> Vec<(String, String)> {
let mut extra_headers: Vec<(String, String)> = vec![
("originator".into(), "zed".into()),
("OpenAI-Beta".into(), "responses=experimental".into()),
];

if let Some(id) = account_id {
if !id.is_empty() {
extra_headers.push(("ChatGPT-Account-Id".into(), id.into()));
}
}

if let Some(id) = session_id {
if !id.is_empty() {
extra_headers.push(("session-id".into(), id.into()));
}
}

extra_headers
}

#[derive(Serialize, Deserialize, Clone, Debug)]
struct CodexCredentials {
Expand Down Expand Up @@ -500,7 +472,6 @@ impl LanguageModel for OpenAiSubscribedLanguageModel {
// The Codex backend rejects `max_output_tokens` (`Unsupported parameter`),
// unlike the public OpenAI Responses API. Pass `None` so the field is
// omitted from the serialized request body entirely.
let session_id = request.thread_id.clone();
let mut responses_request = into_open_ai_response(
request,
self.model.id(),
Expand Down Expand Up @@ -539,24 +510,26 @@ impl LanguageModel for OpenAiSubscribedLanguageModel {
let future = cx.spawn(async move |cx| {
let creds = get_fresh_credentials(&state, &http_client, cx).await?;

let extra_headers =
codex_extra_headers(creds.account_id.as_deref(), session_id.as_deref());
let mut extra_headers: Vec<(String, String)> = vec![
("originator".into(), "zed".into()),
("OpenAI-Beta".into(), "responses=experimental".into()),
];
if let Some(ref id) = creds.account_id {
if !id.is_empty() {
extra_headers.push(("ChatGPT-Account-Id".into(), id.clone()));
}
}

let access_token = creds.access_token.clone();
let background_executor = cx.background_executor().clone();
request_limiter
.stream(async move {
stream_response_with_options(
stream_response(
http_client.as_ref(),
PROVIDER_NAME.0.as_str(),
CODEX_BASE_URL,
&access_token,
responses_request,
extra_headers,
StreamResponseOptions::response_header_timeout(
CODEX_RESPONSE_HEADER_TIMEOUT,
background_executor.timer(CODEX_RESPONSE_HEADER_TIMEOUT),
),
)
.await
.map_err(LanguageModelCompletionError::from)
Expand Down Expand Up @@ -1135,7 +1108,6 @@ mod tests {
use super::*;
use gpui::TestAppContext;
use http_client::FakeHttpClient;
use language_model::{LanguageModelRequestMessage, Role};
use parking_lot::Mutex;
use std::future::Future;
use std::pin::Pin;
Expand Down Expand Up @@ -1185,30 +1157,6 @@ mod tests {
}
}

#[test]
fn test_codex_extra_headers_include_session_id() {
assert_eq!(
codex_extra_headers(Some("account-1"), Some("thread-1")),
vec![
("originator".into(), "zed".into()),
("OpenAI-Beta".into(), "responses=experimental".into()),
("ChatGPT-Account-Id".into(), "account-1".into()),
("session-id".into(), "thread-1".into()),
]
);
}

#[test]
fn test_codex_extra_headers_omit_empty_optional_ids() {
assert_eq!(
codex_extra_headers(Some(""), Some("")),
vec![
("originator".into(), "zed".into()),
("OpenAI-Beta".into(), "responses=experimental".into()),
]
);
}

fn make_expired_credentials() -> CodexCredentials {
CodexCredentials {
access_token: "old_access".to_string(),
Expand All @@ -1229,13 +1177,6 @@ mod tests {
}
}

fn make_fresh_credentials_with_account() -> CodexCredentials {
CodexCredentials {
account_id: Some("account-1".to_string()),
..make_fresh_credentials()
}
}

fn fake_token_response() -> String {
serde_json::json!({
"access_token": "fresh_access",
Expand All @@ -1245,127 +1186,6 @@ mod tests {
.to_string()
}

#[gpui::test]
async fn test_stream_completion_sends_codex_session_header(cx: &mut TestAppContext) {
let captured_headers = Arc::new(Mutex::new(None::<http_client::http::HeaderMap>));
let captured_headers_clone = captured_headers.clone();
let http_client = FakeHttpClient::create(move |request| {
*captured_headers_clone.lock() = Some(request.headers().clone());
async move {
let body = r#"data: {"type":"response.completed","response":{"id":"resp_1","status":"completed"}}"#;
Ok(http_client::Response::builder()
.status(200)
.body(http_client::AsyncBody::from(format!("{body}\n\n")))?)
}
});

let state = cx.new(|_cx| State {
credentials: Some(make_fresh_credentials_with_account()),
sign_in_task: None,
refresh_task: None,
load_task: None,
credentials_provider: Arc::new(FakeCredentialsProvider::new()),
auth_generation: 0,
last_auth_error: None,
});

let model = OpenAiSubscribedLanguageModel {
id: LanguageModelId::from(ChatGptModel::Gpt55.id().to_string()),
model: ChatGptModel::Gpt55,
state,
http_client,
request_limiter: RateLimiter::new(4),
};
let request = LanguageModelRequest {
thread_id: Some("thread-1".to_string()),
prompt_id: Some("prompt-1".to_string()),
messages: vec![LanguageModelRequestMessage {
role: Role::User,
content: vec!["Hello".into()],
cache: false,
reasoning_details: None,
}],
..Default::default()
};

let mut stream = model
.stream_completion(request, &cx.to_async())
.await
.expect("stream should start");
stream
.next()
.await
.expect("stream should emit event")
.expect("event should parse");

let captured_headers = captured_headers
.lock()
.clone()
.expect("request headers should be captured");
assert_eq!(
captured_headers
.get("session-id")
.and_then(|value| value.to_str().ok()),
Some("thread-1")
);
assert_eq!(
captured_headers
.get("ChatGPT-Account-Id")
.and_then(|value| value.to_str().ok()),
Some("account-1")
);
}

#[gpui::test]
async fn test_stream_completion_times_out_before_codex_headers(cx: &mut TestAppContext) {
let http_client = FakeHttpClient::create(|_request| {
futures::future::pending::<anyhow::Result<http_client::Response<AsyncBody>>>()
});

let state = cx.new(|_cx| State {
credentials: Some(make_fresh_credentials()),
sign_in_task: None,
refresh_task: None,
load_task: None,
credentials_provider: Arc::new(FakeCredentialsProvider::new()),
auth_generation: 0,
last_auth_error: None,
});

let model = OpenAiSubscribedLanguageModel {
id: LanguageModelId::from(ChatGptModel::Gpt55.id().to_string()),
model: ChatGptModel::Gpt55,
state,
http_client,
request_limiter: RateLimiter::new(4),
};
let request = LanguageModelRequest {
thread_id: Some("thread-1".to_string()),
prompt_id: Some("prompt-1".to_string()),
messages: vec![LanguageModelRequestMessage {
role: Role::User,
content: vec!["Hello".into()],
cache: false,
reasoning_details: None,
}],
..Default::default()
};

let stream_completion = model.stream_completion(request, &cx.to_async());
cx.run_until_parked();
cx.executor().advance_clock(CODEX_RESPONSE_HEADER_TIMEOUT);

let error = match stream_completion.await {
Ok(_) => panic!("stream should time out before headers arrive"),
Err(error) => error,
};
assert!(matches!(
error,
LanguageModelCompletionError::HttpSend { provider, .. }
if provider == PROVIDER_NAME
));
}

#[gpui::test]
async fn test_concurrent_refresh_deduplicates(cx: &mut TestAppContext) {
let refresh_count = Arc::new(AtomicUsize::new(0));
Expand Down
6 changes: 0 additions & 6 deletions crates/language_models/src/provider/vercel_ai_gateway.rs
Original file line number Diff line number Diff line change
Expand Up @@ -316,12 +316,6 @@ fn map_open_ai_error(error: open_ai::RequestError) -> LanguageModelCompletionErr
retry_after,
)
}
open_ai::RequestError::ResponseHeaderTimeout { timeout, .. } => {
LanguageModelCompletionError::HttpSend {
provider: PROVIDER_NAME,
error: anyhow::anyhow!("response headers timed out after {timeout:?}"),
}
}
open_ai::RequestError::Other(error) => LanguageModelCompletionError::Other(error),
}
}
Expand Down
8 changes: 1 addition & 7 deletions crates/open_ai/src/open_ai.rs
Original file line number Diff line number Diff line change
Expand Up @@ -11,7 +11,7 @@ use http_client::{
pub use language_model_core::ReasoningEffort;
use serde::{Deserialize, Serialize};
use serde_json::Value;
use std::{convert::TryFrom, future::Future, time::Duration};
use std::{convert::TryFrom, future::Future};
use strum::EnumIter;
use thiserror::Error;

Expand Down Expand Up @@ -684,8 +684,6 @@ pub enum RequestError {
body: String,
headers: HeaderMap<HeaderValue>,
},
#[error("response headers from {provider}'s API timed out after {timeout:?}")]
ResponseHeaderTimeout { provider: String, timeout: Duration },
#[error(transparent)]
Other(#[from] anyhow::Error),
}
Expand Down Expand Up @@ -905,10 +903,6 @@ impl From<RequestError> for language_model_core::LanguageModelCompletionError {

Self::from_http_status(provider.into(), status_code, body, retry_after)
}
RequestError::ResponseHeaderTimeout { provider, timeout } => Self::HttpSend {
provider: provider.into(),
error: anyhow!("response headers timed out after {timeout:?}"),
},
RequestError::Other(e) => Self::Other(e),
}
}
Expand Down
Loading
Loading