diff --git a/crates/agent/src/thread.rs b/crates/agent/src/thread.rs index c152082748473a..c34df08273c2f2 100644 --- a/crates/agent/src/thread.rs +++ b/crates/agent/src/thread.rs @@ -4206,10 +4206,8 @@ impl Thread { max_attempts: 3, }) } - Other(err) if err.is::() => { - // 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, diff --git a/crates/agent_ui/src/conversation_view.rs b/crates/agent_ui/src/conversation_view.rs index f59420b1f2b097..95f8b7d80970d9 100644 --- a/crates/agent_ui/src/conversation_view.rs +++ b/crates/agent_ui/src/conversation_view.rs @@ -163,8 +163,6 @@ impl From for ThreadError { Self::MaxOutputTokens } else if error.is::() { Self::NoModelSelected - } else if error.is::() { - Self::PaymentRequired } else if let Some(acp_error) = error.downcast_ref::() && acp_error.code == acp::ErrorCode::AuthRequired { @@ -181,6 +179,7 @@ impl From for ThreadError { } } PromptTooLarge { .. } => Self::PromptTooLarge, + PaymentRequired => Self::PaymentRequired, NoApiKey { provider } => Self::NoApiKey { provider: provider.to_string().into(), }, diff --git a/crates/language_model/src/language_model.rs b/crates/language_model/src/language_model.rs index 5cf1a6ea087b40..1eb2ec5b680f13 100644 --- a/crates/language_model/src/language_model.rs +++ b/crates/language_model/src/language_model.rs @@ -1,5 +1,4 @@ mod api_key; -mod model; mod registry; mod request; @@ -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}; diff --git a/crates/language_model/src/model.rs b/crates/language_model/src/model.rs deleted file mode 100644 index db4c55daa7db99..00000000000000 --- a/crates/language_model/src/model.rs +++ /dev/null @@ -1,3 +0,0 @@ -pub mod cloud_model; - -pub use cloud_model::*; diff --git a/crates/language_model/src/model/cloud_model.rs b/crates/language_model/src/model/cloud_model.rs deleted file mode 100644 index 8cd71928b10fb1..00000000000000 --- a/crates/language_model/src/model/cloud_model.rs +++ /dev/null @@ -1,15 +0,0 @@ -use std::fmt; - -use thiserror::Error; - -#[derive(Error, Debug)] -pub struct PaymentRequiredError; - -impl fmt::Display for PaymentRequiredError { - fn fmt(&self, f: &mut fmt::Formatter) -> fmt::Result { - write!( - f, - "Payment required to use this language model. Please upgrade your account." - ) - } -} diff --git a/crates/language_model_core/src/language_model_core.rs b/crates/language_model_core/src/language_model_core.rs index 3dd8330f71cbb0..bc3d8b8e521e7c 100644 --- a/crates/language_model_core/src/language_model_core.rs +++ b/crates/language_model_core/src/language_model_core.rs @@ -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), } diff --git a/crates/language_models_cloud/src/language_models_cloud.rs b/crates/language_models_cloud/src/language_models_cloud.rs index c129042dec6b51..6fdd11caf571f2 100644 --- a/crates/language_models_cloud/src/language_models_cloud.rs +++ b/crates/language_models_cloud/src/language_models_cloud.rs @@ -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, @@ -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; @@ -123,9 +122,16 @@ impl CloudLanguageModel { auth_context: TP::AuthContext, app_version: Option, body: CompletionBody, - ) -> Result { - let url = http_client.build_zed_llm_url("/completions", &[])?; - let body = serde_json::to_string(&body)?; + ) -> Result { + 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() @@ -140,7 +146,11 @@ impl CloudLanguageModel { .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() { @@ -156,17 +166,25 @@ impl CloudLanguageModel { } 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()) } } @@ -469,15 +487,15 @@ impl LanguageModel for CloudLanguageModel() { - Ok(api_err) => anyhow!(LanguageModelCompletionError::from(api_err)), - Err(err) => anyhow!(err), - })?; + .await?; let mut mapper = AnthropicEventMapper::new(); Ok(map_cloud_completion_events( @@ -534,8 +552,12 @@ impl LanguageModel for CloudLanguageModel LanguageModel for CloudLanguageModel LanguageModel for CloudLanguageModel CloudModelProvider { } pub fn map_cloud_completion_events( - stream: Pin>> + Send>>, + stream: Pin, ResponseStreamError>> + Send>>, provider: &LanguageModelProviderName, mut map_callback: F, ) -> BoxStream<'static, Result> @@ -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; @@ -852,10 +882,36 @@ 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( response: Response, includes_status_messages: bool, -) -> impl Stream>> { +) -> impl Stream, ResponseStreamError>> { futures::stream::try_unfold( (String::new(), BufReader::new(response.into_body())), move |(mut line, mut body)| async move { @@ -863,15 +919,19 @@ pub fn response_lines( Ok(0) => Ok(None), Ok(_) => { let event = if includes_status_messages { - serde_json::from_str::>(&line)? + serde_json::from_str::>(&line) + .map_err(ResponseStreamError::Deserialize)? } else { - CompletionEvent::Event(serde_json::from_str::(&line)?) + CompletionEvent::Event( + serde_json::from_str::(&line) + .map_err(ResponseStreamError::Deserialize)?, + ) }; line.clear(); Ok(Some((event, (line, body)))) } - Err(e) => Err(e.into()), + Err(error) => Err(ResponseStreamError::Read(error)), } }, ) @@ -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::("not json").unwrap_err(), + ) + .into_completion_error(PROVIDER_NAME); + assert!( + matches!( + deserialize, + LanguageModelCompletionError::DeserializeResponse { .. } + ), + "Expected DeserializeResponse, got: {deserialize:?}" + ); + } }