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
6 changes: 2 additions & 4 deletions crates/agent/src/thread.rs
Original file line number Diff line number Diff line change
Expand Up @@ -4206,10 +4206,8 @@ impl Thread {
max_attempts: 3,
})
}
Other(err) if err.is::<language_model::PaymentRequiredError>() => {
// Retrying won't help for Payment Required errors.
None
}
// Retrying won't help for Payment Required errors.
PaymentRequired => None,
// Retrying won't help until the user consents to data retention
// or switches models.
DataRetentionConsentRequired { .. } => None,
Expand Down
3 changes: 1 addition & 2 deletions crates/agent_ui/src/conversation_view.rs
Original file line number Diff line number Diff line change
Expand Up @@ -163,8 +163,6 @@ impl From<anyhow::Error> for ThreadError {
Self::MaxOutputTokens
} else if error.is::<NoModelConfiguredError>() {
Self::NoModelSelected
} else if error.is::<language_model::PaymentRequiredError>() {
Self::PaymentRequired
} else if let Some(acp_error) = error.downcast_ref::<acp::Error>()
&& acp_error.code == acp::ErrorCode::AuthRequired
{
Expand All @@ -181,6 +179,7 @@ impl From<anyhow::Error> for ThreadError {
}
}
PromptTooLarge { .. } => Self::PromptTooLarge,
PaymentRequired => Self::PaymentRequired,
NoApiKey { provider } => Self::NoApiKey {
provider: provider.to_string().into(),
},
Expand Down
2 changes: 0 additions & 2 deletions crates/language_model/src/language_model.rs
Original file line number Diff line number Diff line change
@@ -1,5 +1,4 @@
mod api_key;
mod model;
mod registry;
mod request;

Expand All @@ -17,7 +16,6 @@ use parking_lot::Mutex;
use std::sync::Arc;

pub use crate::api_key::{ApiKey, ApiKeyState};
pub use crate::model::*;
pub use crate::registry::*;
pub use crate::request::{LanguageModelImageExt, gpui_size_to_image_size, image_size_to_gpui};
pub use env_var::{EnvVar, env_var};
Expand Down
3 changes: 0 additions & 3 deletions crates/language_model/src/model.rs

This file was deleted.

15 changes: 0 additions & 15 deletions crates/language_model/src/model/cloud_model.rs

This file was deleted.

2 changes: 2 additions & 0 deletions crates/language_model_core/src/language_model_core.rs
Original file line number Diff line number Diff line change
Expand Up @@ -174,6 +174,8 @@ pub enum LanguageModelCompletionError {
},
#[error("stream from {provider} ended unexpectedly")]
StreamEndedUnexpectedly { provider: LanguageModelProviderName },
#[error("payment required to use this language model; please upgrade your account")]
PaymentRequired,
#[error(transparent)]
Other(#[from] anyhow::Error),
}
Expand Down
152 changes: 120 additions & 32 deletions crates/language_models_cloud/src/language_models_cloud.rs
Original file line number Diff line number Diff line change
@@ -1,5 +1,5 @@
use anthropic::AnthropicModelMode;
use anyhow::{Context as _, Result, anyhow};
use anyhow::{Context as _, Result};
use cloud_llm_client::{
CLIENT_SUPPORTS_STATUS_MESSAGES_HEADER_NAME, CLIENT_SUPPORTS_STATUS_STREAM_ENDED_HEADER_NAME,
CLIENT_SUPPORTS_X_AI_HEADER_NAME, CompletionBody, CompletionEvent, CompletionRequestStatus,
Expand All @@ -23,9 +23,8 @@ use language_model::{
LanguageModel, LanguageModelCompletionError, LanguageModelCompletionEvent,
LanguageModelEffortLevel, LanguageModelId, LanguageModelName, LanguageModelProviderId,
LanguageModelProviderName, LanguageModelRequest, LanguageModelToolChoice,
LanguageModelToolSchemaFormat, OPEN_AI_PROVIDER_ID, OPEN_AI_PROVIDER_NAME,
PaymentRequiredError, RateLimiter, X_AI_PROVIDER_ID, X_AI_PROVIDER_NAME, ZED_CLOUD_PROVIDER_ID,
ZED_CLOUD_PROVIDER_NAME,
LanguageModelToolSchemaFormat, OPEN_AI_PROVIDER_ID, OPEN_AI_PROVIDER_NAME, RateLimiter,
X_AI_PROVIDER_ID, X_AI_PROVIDER_NAME, ZED_CLOUD_PROVIDER_ID, ZED_CLOUD_PROVIDER_NAME,
};

use schemars::JsonSchema;
Expand Down Expand Up @@ -123,9 +122,16 @@ impl<TP: CloudLlmTokenProvider> CloudLanguageModel<TP> {
auth_context: TP::AuthContext,
app_version: Option<Version>,
body: CompletionBody,
) -> Result<PerformLlmCompletionResponse> {
let url = http_client.build_zed_llm_url("/completions", &[])?;
let body = serde_json::to_string(&body)?;
) -> Result<PerformLlmCompletionResponse, LanguageModelCompletionError> {
let url = http_client
.build_zed_llm_url("/completions", &[])
.map_err(LanguageModelCompletionError::Other)?;
let body = serde_json::to_string(&body).map_err(|error| {
LanguageModelCompletionError::SerializeRequest {
provider: PROVIDER_NAME,
error,
}
})?;
let mut response =
authenticated_llm_request(http_client, token_provider, auth_context, |token| {
Ok(http_client::Request::builder()
Expand All @@ -140,7 +146,11 @@ impl<TP: CloudLlmTokenProvider> CloudLanguageModel<TP> {
.header(CLIENT_SUPPORTS_STATUS_STREAM_ENDED_HEADER_NAME, "true")
.body(body.clone().into())?)
})
.await?;
.await
.map_err(|error| LanguageModelCompletionError::HttpSend {
provider: PROVIDER_NAME,
error,
})?;

let status = response.status();
if status.is_success() {
Expand All @@ -156,17 +166,25 @@ impl<TP: CloudLlmTokenProvider> CloudLanguageModel<TP> {
}

if status == StatusCode::PAYMENT_REQUIRED {
return Err(anyhow!(PaymentRequiredError));
return Err(LanguageModelCompletionError::PaymentRequired);
}

let mut body = String::new();
let headers = response.headers().clone();
response.body_mut().read_to_string(&mut body).await?;
Err(anyhow!(ApiError {
response
.body_mut()
.read_to_string(&mut body)
.await
.map_err(|error| LanguageModelCompletionError::ApiReadResponseError {
provider: PROVIDER_NAME,
error,
})?;
Err(ApiError {
status,
body,
headers
}))
headers,
}
.into())
}
}

Expand Down Expand Up @@ -469,15 +487,15 @@ impl<TP: CloudLlmTokenProvider + 'static> LanguageModel for CloudLanguageModel<T
prompt_id,
provider: cloud_llm_client::LanguageModelProvider::Anthropic,
model: request.model.clone(),
provider_request: serde_json::to_value(&request)
.map_err(|e| anyhow!(e))?,
provider_request: serde_json::to_value(&request).map_err(|error| {
LanguageModelCompletionError::SerializeRequest {
provider: provider_name.clone(),
error,
}
})?,
},
)
.await
.map_err(|err| match err.downcast::<ApiError>() {
Ok(api_err) => anyhow!(LanguageModelCompletionError::from(api_err)),
Err(err) => anyhow!(err),
})?;
.await?;

let mut mapper = AnthropicEventMapper::new();
Ok(map_cloud_completion_events(
Expand Down Expand Up @@ -534,8 +552,12 @@ impl<TP: CloudLlmTokenProvider + 'static> LanguageModel for CloudLanguageModel<T
prompt_id,
provider: cloud_llm_client::LanguageModelProvider::OpenAi,
model: request.model.clone(),
provider_request: serde_json::to_value(&request)
.map_err(|e| anyhow!(e))?,
provider_request: serde_json::to_value(&request).map_err(|error| {
LanguageModelCompletionError::SerializeRequest {
provider: provider_name.clone(),
error,
}
})?,
},
)
.await?;
Expand Down Expand Up @@ -576,8 +598,12 @@ impl<TP: CloudLlmTokenProvider + 'static> LanguageModel for CloudLanguageModel<T
prompt_id,
provider: cloud_llm_client::LanguageModelProvider::XAi,
model: request.model.clone(),
provider_request: serde_json::to_value(&request)
.map_err(|e| anyhow!(e))?,
provider_request: serde_json::to_value(&request).map_err(|error| {
LanguageModelCompletionError::SerializeRequest {
provider: provider_name.clone(),
error,
}
})?,
},
)
.await?;
Expand Down Expand Up @@ -611,8 +637,12 @@ impl<TP: CloudLlmTokenProvider + 'static> LanguageModel for CloudLanguageModel<T
prompt_id,
provider: cloud_llm_client::LanguageModelProvider::Google,
model: request.model.model_id.clone(),
provider_request: serde_json::to_value(&request)
.map_err(|e| anyhow!(e))?,
provider_request: serde_json::to_value(&request).map_err(|error| {
LanguageModelCompletionError::SerializeRequest {
provider: provider_name.clone(),
error,
}
})?,
},
)
.await?;
Expand Down Expand Up @@ -772,7 +802,7 @@ impl<TP: CloudLlmTokenProvider + 'static> CloudModelProvider<TP> {
}

pub fn map_cloud_completion_events<T, F>(
stream: Pin<Box<dyn Stream<Item = Result<CompletionEvent<T>>> + Send>>,
stream: Pin<Box<dyn Stream<Item = Result<CompletionEvent<T>, ResponseStreamError>> + Send>>,
provider: &LanguageModelProviderName,
mut map_callback: F,
) -> BoxStream<'static, Result<LanguageModelCompletionEvent, LanguageModelCompletionError>>
Expand Down Expand Up @@ -804,7 +834,7 @@ where
Poll::Ready(Some(event)) => {
let items = match event {
Err(error) => {
vec![Err(LanguageModelCompletionError::from(error))]
vec![Err(error.into_completion_error(provider.clone()))]
}
Ok(CompletionEvent::Status(CompletionRequestStatus::StreamEnded)) => {
saw_stream_ended = true;
Expand Down Expand Up @@ -852,26 +882,56 @@ pub fn provider_name(
}
}

/// A failure while reading the streamed completion response body.
///
/// Kept as a typed error (rather than `anyhow::Error`) so the consumer can
/// attach the provider name and build a structured
/// [`LanguageModelCompletionError`] without a runtime downcast.
pub enum ResponseStreamError {
Read(std::io::Error),
Deserialize(serde_json::Error),
}

impl ResponseStreamError {
fn into_completion_error(
self,
provider: LanguageModelProviderName,
) -> LanguageModelCompletionError {
match self {
ResponseStreamError::Read(error) => {
LanguageModelCompletionError::ApiReadResponseError { provider, error }
}
ResponseStreamError::Deserialize(error) => {
LanguageModelCompletionError::DeserializeResponse { provider, error }
}
}
}
}

pub fn response_lines<T: DeserializeOwned>(
response: Response<AsyncBody>,
includes_status_messages: bool,
) -> impl Stream<Item = Result<CompletionEvent<T>>> {
) -> impl Stream<Item = Result<CompletionEvent<T>, ResponseStreamError>> {
futures::stream::try_unfold(
(String::new(), BufReader::new(response.into_body())),
move |(mut line, mut body)| async move {
match body.read_line(&mut line).await {
Ok(0) => Ok(None),
Ok(_) => {
let event = if includes_status_messages {
serde_json::from_str::<CompletionEvent<T>>(&line)?
serde_json::from_str::<CompletionEvent<T>>(&line)
.map_err(ResponseStreamError::Deserialize)?
} else {
CompletionEvent::Event(serde_json::from_str::<T>(&line)?)
CompletionEvent::Event(
serde_json::from_str::<T>(&line)
.map_err(ResponseStreamError::Deserialize)?,
)
};

line.clear();
Ok(Some((event, (line, body))))
}
Err(e) => Err(e.into()),
Err(error) => Err(ResponseStreamError::Read(error)),
}
},
)
Expand Down Expand Up @@ -1027,4 +1087,32 @@ mod tests {
),
}
}

#[test]
fn test_response_stream_error_maps_to_structured_variant() {
// Read/deserialize failures mid-stream must keep their structured
// variant rather than collapsing into `Other` (the source of the
// generic "Request failed." message).
let read = ResponseStreamError::Read(std::io::Error::from(std::io::ErrorKind::BrokenPipe))
.into_completion_error(PROVIDER_NAME);
assert!(
matches!(
read,
LanguageModelCompletionError::ApiReadResponseError { .. }
),
"Expected ApiReadResponseError, got: {read:?}"
);

let deserialize = ResponseStreamError::Deserialize(
serde_json::from_str::<serde_json::Value>("not json").unwrap_err(),
)
.into_completion_error(PROVIDER_NAME);
assert!(
matches!(
deserialize,
LanguageModelCompletionError::DeserializeResponse { .. }
),
"Expected DeserializeResponse, got: {deserialize:?}"
);
}
}
Loading