Skip to repository content

tenant.openagents/omega

No repository description is available.

OpenAgents Git authority 2026-07-28T02:32:55.341Z Public web read
NIP-34 coordinate30617:7649603503856e5148d571eac2766b288a8ff1e9e35d380337a1d2b0015b4f92:omega
MaintainersHidden in public view
References2 branches · 1 tag
Read-only clonegit clone https://openagents.com/git/tenant.openagents/omega.git
Browse files

language_models_cloud.rs

1530 lines · 58.6 KB · rust
1use anthropic::AnthropicModelMode;
2use anyhow::{Context as _, Result};
3use cloud_llm_client::{
4    CLIENT_SUPPORTS_STATUS_MESSAGES_HEADER_NAME, CLIENT_SUPPORTS_STATUS_STREAM_ENDED_HEADER_NAME,
5    CLIENT_SUPPORTS_X_AI_HEADER_NAME, CompletionBody, CompletionEvent, CompletionRequestStatus,
6    EXPIRED_LLM_TOKEN_HEADER_NAME, ListModelsResponse, OUTDATED_LLM_TOKEN_HEADER_NAME,
7    SERVER_SUPPORTS_STATUS_MESSAGES_HEADER_NAME, ZED_VERSION_HEADER_NAME,
8};
9use futures::{
10    AsyncBufReadExt, AsyncReadExt as _, FutureExt, Stream, StreamExt,
11    future::BoxFuture,
12    io::BufReader,
13    stream::{self, BoxStream},
14};
15use google_ai::GoogleModelMode;
16use gpui::{AppContext, AsyncApp, Context, Task};
17use http_client::http::{HeaderMap, HeaderValue};
18use http_client::{
19    AsyncBody, HttpClient, HttpClientWithUrl, HttpRequestExt, Method, Response, StatusCode,
20};
21use language_model::{
22    ANTHROPIC_PROVIDER_ID, ANTHROPIC_PROVIDER_NAME, CompactionResult, DisabledReason,
23    GOOGLE_PROVIDER_ID, GOOGLE_PROVIDER_NAME, LanguageModel, LanguageModelCompletionError,
24    LanguageModelCompletionEvent, LanguageModelEffortLevel, LanguageModelId, LanguageModelName,
25    LanguageModelProviderId, LanguageModelProviderName, LanguageModelRequest,
26    LanguageModelToolChoice, LanguageModelToolSchemaFormat, OPEN_AI_PROVIDER_ID,
27    OPEN_AI_PROVIDER_NAME, RateLimiter, X_AI_PROVIDER_ID, X_AI_PROVIDER_NAME,
28    ZED_CLOUD_PROVIDER_ID, ZED_CLOUD_PROVIDER_NAME,
29};
30
31use schemars::JsonSchema;
32use semver::Version;
33use serde::{Deserialize, Serialize, de::DeserializeOwned};
34use std::collections::VecDeque;
35use std::pin::Pin;
36use std::str::FromStr;
37use std::sync::Arc;
38use std::task::Poll;
39use std::time::Duration;
40use thiserror::Error;
41
42use anthropic::completion::{AnthropicEventMapper, AnthropicPromptCacheMode, into_anthropic};
43use google_ai::completion::{GoogleEventMapper, into_google};
44use open_ai::completion::{
45    ChatCompletionMaxTokensParameter, OpenAiEventMapper, OpenAiResponseEventMapper, into_open_ai,
46    into_open_ai_response, token_usage_from_response_usage,
47};
48
49const PROVIDER_ID: LanguageModelProviderId = ZED_CLOUD_PROVIDER_ID;
50const PROVIDER_NAME: LanguageModelProviderName = ZED_CLOUD_PROVIDER_NAME;
51
52/// Trait for acquiring and refreshing LLM authentication tokens.
53pub trait CloudLlmTokenProvider: Send + Sync {
54    type AuthContext: Clone + Send + 'static;
55
56    fn auth_context(&self, cx: &impl AppContext) -> Self::AuthContext;
57    fn cached_token(&self, auth_context: Self::AuthContext) -> BoxFuture<'static, Result<String>>;
58    fn refresh_token(&self, auth_context: Self::AuthContext) -> BoxFuture<'static, Result<String>>;
59
60    /// Whether the user has consented to upstream providers retaining
61    /// inference logs for models that require it (see
62    /// [`LanguageModel::requires_data_retention`]).
63    fn has_data_retention_consent(&self, cx: &impl AppContext) -> bool;
64}
65
66/// Sends an authenticated request to the Zed LLM service, retrying once with
67/// a refreshed token if the server signals that the cached LLM token is
68/// expired or otherwise rejected. Returns the raw response so callers can
69/// inspect headers and stream the body.
70pub async fn authenticated_llm_request<TP: CloudLlmTokenProvider>(
71    http_client: &HttpClientWithUrl,
72    token_provider: &TP,
73    auth_context: TP::AuthContext,
74    build_request: impl Fn(&str) -> Result<http_client::Request<AsyncBody>>,
75) -> Result<Response<AsyncBody>> {
76    let token = token_provider.cached_token(auth_context.clone()).await?;
77    let response = http_client.send(build_request(&token)?).await?;
78    if !needs_llm_token_refresh(&response) && response.status() != StatusCode::UNAUTHORIZED {
79        return Ok(response);
80    }
81    log::info!("LLM token rejected; refreshing and retrying request");
82    let token = token_provider.refresh_token(auth_context).await?;
83    http_client.send(build_request(&token)?).await
84}
85
86#[derive(Default, Clone, Debug, PartialEq, Serialize, Deserialize, JsonSchema)]
87#[serde(tag = "type", rename_all = "lowercase")]
88pub enum ModelMode {
89    #[default]
90    Default,
91    Thinking {
92        /// The maximum number of tokens to use for reasoning. Must be lower than the model's `max_output_tokens`.
93        budget_tokens: Option<u32>,
94    },
95}
96
97impl From<ModelMode> for AnthropicModelMode {
98    fn from(value: ModelMode) -> Self {
99        match value {
100            ModelMode::Default => AnthropicModelMode::Default,
101            ModelMode::Thinking { budget_tokens } => AnthropicModelMode::Thinking { budget_tokens },
102        }
103    }
104}
105
106pub struct CloudLanguageModel<TP: CloudLlmTokenProvider> {
107    pub id: LanguageModelId,
108    pub model: Arc<cloud_llm_client::LanguageModel>,
109    pub token_provider: Arc<TP>,
110    pub http_client: Arc<HttpClientWithUrl>,
111    pub app_version: Option<Version>,
112    pub request_limiter: RateLimiter,
113}
114
115pub struct PerformLlmCompletionResponse {
116    pub response: Response<AsyncBody>,
117    pub includes_status_messages: bool,
118}
119
120impl<TP: CloudLlmTokenProvider> CloudLanguageModel<TP> {
121    pub async fn perform_llm_completion(
122        http_client: &HttpClientWithUrl,
123        token_provider: &TP,
124        auth_context: TP::AuthContext,
125        app_version: Option<Version>,
126        body: CompletionBody,
127    ) -> Result<PerformLlmCompletionResponse, LanguageModelCompletionError> {
128        Self::perform_llm_request(
129            "/completions",
130            true,
131            http_client,
132            token_provider,
133            auth_context,
134            app_version,
135            body,
136        )
137        .await
138    }
139
140    async fn perform_llm_compaction(
141        http_client: &HttpClientWithUrl,
142        token_provider: &TP,
143        auth_context: TP::AuthContext,
144        app_version: Option<Version>,
145        body: CompletionBody,
146    ) -> Result<PerformLlmCompletionResponse, LanguageModelCompletionError> {
147        Self::perform_llm_request(
148            "/completions/compact",
149            false,
150            http_client,
151            token_provider,
152            auth_context,
153            app_version,
154            body,
155        )
156        .await
157    }
158
159    async fn perform_llm_request(
160        path: &str,
161        request_status_messages: bool,
162        http_client: &HttpClientWithUrl,
163        token_provider: &TP,
164        auth_context: TP::AuthContext,
165        app_version: Option<Version>,
166        body: CompletionBody,
167    ) -> Result<PerformLlmCompletionResponse, LanguageModelCompletionError> {
168        let url = http_client
169            .build_zed_llm_url(path, &[])
170            .map_err(LanguageModelCompletionError::Other)?;
171        let body = serde_json::to_string(&body).map_err(|error| {
172            LanguageModelCompletionError::SerializeRequest {
173                provider: PROVIDER_NAME,
174                error,
175            }
176        })?;
177        let mut response =
178            authenticated_llm_request(http_client, token_provider, auth_context, |token| {
179                let mut request = http_client::Request::builder()
180                    .method(Method::POST)
181                    .uri(url.as_ref())
182                    .when_some(app_version.as_ref(), |builder, app_version| {
183                        builder.header(ZED_VERSION_HEADER_NAME, app_version.to_string())
184                    })
185                    .header("Content-Type", "application/json")
186                    .header("Authorization", format!("Bearer {token}"));
187                if request_status_messages {
188                    request = request
189                        .header(CLIENT_SUPPORTS_STATUS_MESSAGES_HEADER_NAME, "true")
190                        .header(CLIENT_SUPPORTS_STATUS_STREAM_ENDED_HEADER_NAME, "true");
191                }
192                Ok(request.body(body.clone().into())?)
193            })
194            .await
195            .map_err(|error| LanguageModelCompletionError::HttpSend {
196                provider: PROVIDER_NAME,
197                error,
198            })?;
199
200        let status = response.status();
201        if status.is_success() {
202            let includes_status_messages = request_status_messages
203                && response
204                    .headers()
205                    .get(SERVER_SUPPORTS_STATUS_MESSAGES_HEADER_NAME)
206                    .is_some();
207
208            return Ok(PerformLlmCompletionResponse {
209                response,
210                includes_status_messages,
211            });
212        }
213
214        if status == StatusCode::PAYMENT_REQUIRED {
215            return Err(LanguageModelCompletionError::PaymentRequired);
216        }
217
218        let mut body = String::new();
219        let headers = response.headers().clone();
220        response
221            .body_mut()
222            .read_to_string(&mut body)
223            .await
224            .map_err(|error| LanguageModelCompletionError::ApiReadResponseError {
225                provider: PROVIDER_NAME,
226                error,
227            })?;
228        Err(ApiError {
229            status,
230            body,
231            headers,
232        }
233        .into())
234    }
235}
236
237fn needs_llm_token_refresh(response: &Response<AsyncBody>) -> bool {
238    response
239        .headers()
240        .get(EXPIRED_LLM_TOKEN_HEADER_NAME)
241        .is_some()
242        || response
243            .headers()
244            .get(OUTDATED_LLM_TOKEN_HEADER_NAME)
245            .is_some()
246}
247
248#[derive(Debug, Error)]
249#[error("cloud language model request failed with status {status}: {body}")]
250struct ApiError {
251    status: StatusCode,
252    body: String,
253    headers: HeaderMap<HeaderValue>,
254}
255
256/// Represents error responses from Zed's cloud API.
257///
258/// Example JSON for an upstream HTTP error:
259/// ```json
260/// {
261///   "code": "upstream_http_error",
262///   "message": "Received an error from the Anthropic API: upstream connect error or disconnect/reset before headers, reset reason: connection timeout",
263///   "upstream_status": 503
264/// }
265/// ```
266#[derive(Debug, serde::Deserialize)]
267struct CloudApiError {
268    code: String,
269    message: String,
270    #[serde(default)]
271    #[serde(deserialize_with = "deserialize_optional_status_code")]
272    upstream_status: Option<StatusCode>,
273    #[serde(default)]
274    retry_after: Option<f64>,
275}
276
277fn deserialize_optional_status_code<'de, D>(deserializer: D) -> Result<Option<StatusCode>, D::Error>
278where
279    D: serde::Deserializer<'de>,
280{
281    let opt: Option<u16> = Option::deserialize(deserializer)?;
282    Ok(opt.and_then(|code| StatusCode::from_u16(code).ok()))
283}
284
285impl From<ApiError> for LanguageModelCompletionError {
286    fn from(error: ApiError) -> Self {
287        if let Ok(cloud_error) = serde_json::from_str::<CloudApiError>(&error.body) {
288            if cloud_error.code.starts_with("upstream_http_") {
289                let status = if let Some(status) = cloud_error.upstream_status {
290                    status
291                } else if cloud_error.code.ends_with("_error") {
292                    error.status
293                } else {
294                    // If there's a status code in the code string (e.g. "upstream_http_429")
295                    // then use that; otherwise, see if the JSON contains a status code.
296                    cloud_error
297                        .code
298                        .strip_prefix("upstream_http_")
299                        .and_then(|code_str| code_str.parse::<u16>().ok())
300                        .and_then(|code| StatusCode::from_u16(code).ok())
301                        .unwrap_or(error.status)
302                };
303
304                return LanguageModelCompletionError::UpstreamProviderError {
305                    message: cloud_error.message,
306                    status,
307                    retry_after: cloud_error.retry_after.map(Duration::from_secs_f64),
308                };
309            }
310
311            return LanguageModelCompletionError::from_http_status(
312                PROVIDER_NAME,
313                error.status,
314                cloud_error.message,
315                None,
316            );
317        }
318
319        let retry_after = None;
320        LanguageModelCompletionError::from_http_status(
321            PROVIDER_NAME,
322            error.status,
323            error.body,
324            retry_after,
325        )
326    }
327}
328
329impl<TP: CloudLlmTokenProvider + 'static> LanguageModel for CloudLanguageModel<TP> {
330    fn id(&self) -> LanguageModelId {
331        self.id.clone()
332    }
333
334    fn name(&self) -> LanguageModelName {
335        LanguageModelName::from(self.model.display_name.clone())
336    }
337
338    fn provider_id(&self) -> LanguageModelProviderId {
339        PROVIDER_ID
340    }
341
342    fn provider_name(&self) -> LanguageModelProviderName {
343        PROVIDER_NAME
344    }
345
346    fn upstream_provider_id(&self) -> LanguageModelProviderId {
347        use cloud_llm_client::LanguageModelProvider::*;
348        match self.model.provider {
349            Anthropic => ANTHROPIC_PROVIDER_ID,
350            OpenAi => OPEN_AI_PROVIDER_ID,
351            Google => GOOGLE_PROVIDER_ID,
352            XAi => X_AI_PROVIDER_ID,
353        }
354    }
355
356    fn upstream_provider_name(&self) -> LanguageModelProviderName {
357        use cloud_llm_client::LanguageModelProvider::*;
358        match self.model.provider {
359            Anthropic => ANTHROPIC_PROVIDER_NAME,
360            OpenAi => OPEN_AI_PROVIDER_NAME,
361            Google => GOOGLE_PROVIDER_NAME,
362            XAi => X_AI_PROVIDER_NAME,
363        }
364    }
365
366    fn is_latest(&self) -> bool {
367        self.model.is_latest
368    }
369
370    fn is_disabled(&self) -> Option<DisabledReason> {
371        if self.model.is_disabled {
372            self.model.disabled_reason.clone().map(DisabledReason::new)
373        } else {
374            None
375        }
376    }
377
378    fn requires_data_retention(&self) -> bool {
379        // Anthropic cannot offer Fable models with Zero Data Retention
380        self.id
381            .0
382            .as_ref()
383            .starts_with(anthropic::FABLE_MODEL_ID_PREFIX)
384    }
385
386    fn refusal_fallback_model_id(&self) -> Option<&'static str> {
387        if self
388            .id
389            .0
390            .as_ref()
391            .starts_with(anthropic::FABLE_MODEL_ID_PREFIX)
392        {
393            Some(anthropic::FABLE_FALLBACK_MODEL_ID)
394        } else {
395            None
396        }
397    }
398
399    fn supports_tools(&self) -> bool {
400        self.model.supports_tools
401    }
402
403    fn supports_images(&self) -> bool {
404        self.model.supports_images
405    }
406
407    fn supports_thinking(&self) -> bool {
408        self.model.supports_thinking
409    }
410
411    fn supports_disabling_thinking(&self) -> bool {
412        self.model.supports_disabling_thinking
413    }
414
415    fn supports_fast_mode(&self) -> bool {
416        self.model.supports_fast_mode
417    }
418
419    fn supports_server_side_compaction(&self) -> bool {
420        self.model.supports_server_side_compaction
421    }
422
423    fn supports_explicit_compaction(&self) -> bool {
424        self.model.provider == cloud_llm_client::LanguageModelProvider::OpenAi
425            && self.model.supports_server_side_compaction
426    }
427
428    fn compact(
429        &self,
430        request: LanguageModelRequest,
431        cx: &AsyncApp,
432    ) -> BoxFuture<'static, Result<CompactionResult, LanguageModelCompletionError>> {
433        if !self.supports_explicit_compaction() {
434            return async {
435                Err(LanguageModelCompletionError::Other(anyhow::anyhow!(
436                    "this cloud model does not support explicit compaction"
437                )))
438            }
439            .boxed();
440        }
441
442        let thread_id = request.thread_id.clone();
443        let prompt_id = request.prompt_id.clone();
444        let app_version = self.app_version.clone();
445        let model_provider = self.model.provider;
446        let provider_name = provider_name(&self.model.provider);
447        let supports_none_reasoning_effort =
448            self.model.supported_effort_levels.iter().any(|effort| {
449                open_ai::ReasoningEffort::from_str(&effort.value)
450                    .is_ok_and(|effort| effort == open_ai::ReasoningEffort::None)
451            });
452        // Cloud proxies to OpenAI's own infrastructure, so the resulting
453        // compaction state is owned by (and interchangeable with) OpenAI
454        // proper, not by the cloud transport.
455        let request = match into_open_ai_response(
456            request,
457            &self.model.id.0,
458            self.model.supports_parallel_tool_calls,
459            true,
460            None,
461            None,
462            supports_none_reasoning_effort,
463            &OPEN_AI_PROVIDER_ID,
464        ) {
465            Ok(request) => request,
466            Err(error) => return async move { Err(error.into()) }.boxed(),
467        };
468        let compact_request = request.into_compact_request();
469        let http_client = self.http_client.clone();
470        let token_provider = self.token_provider.clone();
471        let auth_context = token_provider.auth_context(cx);
472        let future = self.request_limiter.run(async move {
473            let PerformLlmCompletionResponse {
474                response,
475                includes_status_messages,
476            } = Self::perform_llm_compaction(
477                &http_client,
478                &*token_provider,
479                auth_context,
480                app_version,
481                CompletionBody {
482                    thread_id,
483                    prompt_id,
484                    provider: model_provider,
485                    model: compact_request.model.clone(),
486                    provider_request: serde_json::to_value(compact_request).map_err(|error| {
487                        LanguageModelCompletionError::SerializeRequest {
488                            provider: provider_name.clone(),
489                            error,
490                        }
491                    })?,
492                },
493            )
494            .await?;
495
496            let events = response_lines::<open_ai::responses::CompactedResponse>(
497                response,
498                includes_status_messages,
499            );
500            futures::pin_mut!(events);
501            while let Some(event) = events.next().await {
502                match event.map_err(|error| error.into_completion_error(provider_name.clone()))? {
503                    CompletionEvent::Event(response) => {
504                        let usage = token_usage_from_response_usage(&response.usage);
505                        let context = response
506                            .into_compacted_context(OPEN_AI_PROVIDER_ID)
507                            .map_err(LanguageModelCompletionError::Other)?;
508                        return Ok(CompactionResult { context, usage });
509                    }
510                    CompletionEvent::Status(_) => {}
511                }
512            }
513
514            Err(LanguageModelCompletionError::StreamEndedUnexpectedly {
515                provider: provider_name,
516            })
517        });
518        future.boxed()
519    }
520
521    fn supported_effort_levels(&self) -> Vec<LanguageModelEffortLevel> {
522        self.model
523            .supported_effort_levels
524            .iter()
525            .map(|effort_level| LanguageModelEffortLevel {
526                name: effort_level.name.clone().into(),
527                value: effort_level.value.clone().into(),
528                is_default: effort_level.is_default.unwrap_or(false),
529            })
530            .collect()
531    }
532
533    fn supports_streaming_tools(&self) -> bool {
534        self.model.supports_streaming_tools
535    }
536
537    fn supports_tool_choice(&self, choice: LanguageModelToolChoice) -> bool {
538        match choice {
539            LanguageModelToolChoice::Auto
540            | LanguageModelToolChoice::Any
541            | LanguageModelToolChoice::None => true,
542        }
543    }
544
545    fn supports_split_token_display(&self) -> bool {
546        use cloud_llm_client::LanguageModelProvider::*;
547        matches!(self.model.provider, OpenAi | XAi)
548    }
549
550    fn telemetry_id(&self) -> String {
551        format!("zed.dev/{}", self.model.id)
552    }
553
554    fn tool_input_format(&self) -> LanguageModelToolSchemaFormat {
555        match self.model.provider {
556            cloud_llm_client::LanguageModelProvider::Anthropic
557            | cloud_llm_client::LanguageModelProvider::OpenAi => {
558                LanguageModelToolSchemaFormat::JsonSchema
559            }
560            cloud_llm_client::LanguageModelProvider::Google
561            | cloud_llm_client::LanguageModelProvider::XAi => {
562                LanguageModelToolSchemaFormat::JsonSchemaSubset
563            }
564        }
565    }
566
567    fn max_token_count(&self) -> u64 {
568        self.model.max_token_count as u64
569    }
570
571    fn max_output_tokens(&self) -> Option<u64> {
572        Some(self.model.max_output_tokens as u64)
573    }
574
575    fn stream_completion(
576        &self,
577        request: LanguageModelRequest,
578        cx: &AsyncApp,
579    ) -> BoxFuture<
580        'static,
581        Result<
582            BoxStream<'static, Result<LanguageModelCompletionEvent, LanguageModelCompletionError>>,
583            LanguageModelCompletionError,
584        >,
585    > {
586        if self.requires_data_retention() && !self.token_provider.has_data_retention_consent(cx) {
587            let model_name = self.model.display_name.clone();
588            return async move {
589                Err(LanguageModelCompletionError::DataRetentionConsentRequired { model_name })
590            }
591            .boxed();
592        }
593
594        let thread_id = request.thread_id.clone();
595        let prompt_id = request.prompt_id.clone();
596        let app_version = self.app_version.clone();
597        let thinking_allowed = request.thinking_allowed;
598        let enable_thinking = thinking_allowed && self.model.supports_thinking;
599        let provider_name = provider_name(&self.model.provider);
600        match self.model.provider {
601            cloud_llm_client::LanguageModelProvider::Anthropic => {
602                let effort = request
603                    .thinking_effort
604                    .as_ref()
605                    .and_then(|effort| anthropic::Effort::from_str(effort).ok());
606
607                let mut request = match into_anthropic(
608                    request,
609                    self.model.id.to_string(),
610                    1.0,
611                    self.model.max_output_tokens as u64,
612                    if enable_thinking {
613                        AnthropicModelMode::Thinking {
614                            budget_tokens: Some(4_096),
615                        }
616                    } else {
617                        AnthropicModelMode::Default
618                    },
619                    AnthropicPromptCacheMode::Automatic,
620                    // Cloud proxies to Anthropic's own infrastructure, so
621                    // compaction state is owned by (and interchangeable with)
622                    // Anthropic proper, not by the cloud transport.
623                    &ANTHROPIC_PROVIDER_ID,
624                ) {
625                    Ok(request) => request,
626                    Err(error) => return async move { Err(error.into()) }.boxed(),
627                };
628
629                if enable_thinking && effort.is_some() {
630                    request.thinking = Some(anthropic::Thinking::Adaptive {
631                        display: Some(anthropic::AdaptiveThinkingDisplay::Summarized),
632                    });
633                    request.output_config = Some(anthropic::OutputConfig { effort });
634                }
635
636                if !self.model.supports_fast_mode {
637                    request.speed = None;
638                }
639
640                let http_client = self.http_client.clone();
641                let token_provider = self.token_provider.clone();
642                let auth_context = token_provider.auth_context(cx);
643                let future = self.request_limiter.stream(async move {
644                    let PerformLlmCompletionResponse {
645                        response,
646                        includes_status_messages,
647                    } = Self::perform_llm_completion(
648                        &http_client,
649                        &*token_provider,
650                        auth_context,
651                        app_version,
652                        CompletionBody {
653                            thread_id,
654                            prompt_id,
655                            provider: cloud_llm_client::LanguageModelProvider::Anthropic,
656                            model: request.model.clone(),
657                            provider_request: serde_json::to_value(&request).map_err(|error| {
658                                LanguageModelCompletionError::SerializeRequest {
659                                    provider: provider_name.clone(),
660                                    error,
661                                }
662                            })?,
663                        },
664                    )
665                    .await?;
666
667                    let mut mapper =
668                        AnthropicEventMapper::new(provider_name.clone(), ANTHROPIC_PROVIDER_ID);
669                    Ok(map_cloud_completion_events(
670                        Box::pin(response_lines(response, includes_status_messages)),
671                        &provider_name,
672                        move |event| mapper.map_event(event),
673                    ))
674                });
675                async move { Ok(future.await?.boxed()) }.boxed()
676            }
677            cloud_llm_client::LanguageModelProvider::OpenAi => {
678                let http_client = self.http_client.clone();
679                let token_provider = self.token_provider.clone();
680                let effort = request
681                    .thinking_effort
682                    .as_ref()
683                    .and_then(|effort| open_ai::ReasoningEffort::from_str(effort).ok())
684                    .filter(|effort| *effort != open_ai::ReasoningEffort::None);
685                let supports_none_reasoning_effort =
686                    self.model.supported_effort_levels.iter().any(|effort| {
687                        open_ai::ReasoningEffort::from_str(&effort.value)
688                            .is_ok_and(|effort| effort == open_ai::ReasoningEffort::None)
689                    });
690
691                let mut request = match into_open_ai_response(
692                    request,
693                    &self.model.id.0,
694                    self.model.supports_parallel_tool_calls,
695                    true,
696                    None,
697                    None,
698                    supports_none_reasoning_effort,
699                    &OPEN_AI_PROVIDER_ID,
700                ) {
701                    Ok(request) => request,
702                    Err(error) => return async move { Err(error.into()) }.boxed(),
703                };
704
705                if enable_thinking && let Some(effort) = effort {
706                    request.reasoning = Some(open_ai::responses::ReasoningConfig {
707                        effort,
708                        summary: Some(open_ai::responses::ReasoningSummaryMode::Auto),
709                    });
710                }
711
712                let auth_context = token_provider.auth_context(cx);
713                let future = self.request_limiter.stream(async move {
714                    let PerformLlmCompletionResponse {
715                        response,
716                        includes_status_messages,
717                    } = Self::perform_llm_completion(
718                        &http_client,
719                        &*token_provider,
720                        auth_context,
721                        app_version,
722                        CompletionBody {
723                            thread_id,
724                            prompt_id,
725                            provider: cloud_llm_client::LanguageModelProvider::OpenAi,
726                            model: request.model.clone(),
727                            provider_request: serde_json::to_value(&request).map_err(|error| {
728                                LanguageModelCompletionError::SerializeRequest {
729                                    provider: provider_name.clone(),
730                                    error,
731                                }
732                            })?,
733                        },
734                    )
735                    .await?;
736
737                    let mut mapper = OpenAiResponseEventMapper::new(OPEN_AI_PROVIDER_ID);
738                    Ok(map_cloud_completion_events(
739                        Box::pin(response_lines(response, includes_status_messages)),
740                        &provider_name,
741                        move |event| mapper.map_event(event),
742                    ))
743                });
744                async move { Ok(future.await?.boxed()) }.boxed()
745            }
746            cloud_llm_client::LanguageModelProvider::XAi => {
747                let http_client = self.http_client.clone();
748                let token_provider = self.token_provider.clone();
749                let request = match into_open_ai(
750                    request,
751                    &self.model.id.0,
752                    self.model.supports_parallel_tool_calls,
753                    false,
754                    None,
755                    ChatCompletionMaxTokensParameter::MaxCompletionTokens,
756                    None,
757                    false,
758                ) {
759                    Ok(request) => request,
760                    Err(error) => return async move { Err(error.into()) }.boxed(),
761                };
762                let auth_context = token_provider.auth_context(cx);
763                let future = self.request_limiter.stream(async move {
764                    let PerformLlmCompletionResponse {
765                        response,
766                        includes_status_messages,
767                    } = Self::perform_llm_completion(
768                        &http_client,
769                        &*token_provider,
770                        auth_context,
771                        app_version,
772                        CompletionBody {
773                            thread_id,
774                            prompt_id,
775                            provider: cloud_llm_client::LanguageModelProvider::XAi,
776                            model: request.model.clone(),
777                            provider_request: serde_json::to_value(&request).map_err(|error| {
778                                LanguageModelCompletionError::SerializeRequest {
779                                    provider: provider_name.clone(),
780                                    error,
781                                }
782                            })?,
783                        },
784                    )
785                    .await?;
786
787                    let mut mapper = OpenAiEventMapper::new();
788                    Ok(map_cloud_completion_events(
789                        Box::pin(response_lines(response, includes_status_messages)),
790                        &provider_name,
791                        move |event| mapper.map_event(event),
792                    ))
793                });
794                async move { Ok(future.await?.boxed()) }.boxed()
795            }
796            cloud_llm_client::LanguageModelProvider::Google => {
797                let http_client = self.http_client.clone();
798                let token_provider = self.token_provider.clone();
799                let request =
800                    match into_google(request, self.model.id.to_string(), GoogleModelMode::Default)
801                    {
802                        Ok(request) => request,
803                        Err(error) => return async move { Err(error.into()) }.boxed(),
804                    };
805                let auth_context = token_provider.auth_context(cx);
806                let future = self.request_limiter.stream(async move {
807                    let PerformLlmCompletionResponse {
808                        response,
809                        includes_status_messages,
810                    } = Self::perform_llm_completion(
811                        &http_client,
812                        &*token_provider,
813                        auth_context,
814                        app_version,
815                        CompletionBody {
816                            thread_id,
817                            prompt_id,
818                            provider: cloud_llm_client::LanguageModelProvider::Google,
819                            model: request.model.model_id.clone(),
820                            provider_request: serde_json::to_value(&request).map_err(|error| {
821                                LanguageModelCompletionError::SerializeRequest {
822                                    provider: provider_name.clone(),
823                                    error,
824                                }
825                            })?,
826                        },
827                    )
828                    .await?;
829
830                    let mut mapper = GoogleEventMapper::new();
831                    Ok(map_cloud_completion_events(
832                        Box::pin(response_lines(response, includes_status_messages)),
833                        &provider_name,
834                        move |event| mapper.map_event(event),
835                    ))
836                });
837                async move { Ok(future.await?.boxed()) }.boxed()
838            }
839        }
840    }
841}
842
843pub struct CloudModelProvider<TP: CloudLlmTokenProvider> {
844    token_provider: Arc<TP>,
845    http_client: Arc<HttpClientWithUrl>,
846    app_version: Option<Version>,
847    models: Vec<Arc<cloud_llm_client::LanguageModel>>,
848    default_model: Option<Arc<cloud_llm_client::LanguageModel>>,
849    default_fast_model: Option<Arc<cloud_llm_client::LanguageModel>>,
850    recommended_models: Vec<Arc<cloud_llm_client::LanguageModel>>,
851}
852
853impl<TP: CloudLlmTokenProvider + 'static> CloudModelProvider<TP> {
854    pub fn new(
855        token_provider: Arc<TP>,
856        http_client: Arc<HttpClientWithUrl>,
857        app_version: Option<Version>,
858    ) -> Self {
859        Self {
860            token_provider,
861            http_client,
862            app_version,
863            models: Vec::new(),
864            default_model: None,
865            default_fast_model: None,
866            recommended_models: Vec::new(),
867        }
868    }
869
870    pub fn refresh_models(&self, cx: &mut Context<Self>) -> Task<Result<()>> {
871        let http_client = self.http_client.clone();
872        let token_provider = self.token_provider.clone();
873        cx.spawn(async move |this, cx| {
874            let auth_context = token_provider.auth_context(cx);
875            let response =
876                Self::fetch_models_request(&http_client, &*token_provider, auth_context).await?;
877            this.update(cx, |this, cx| {
878                this.update_models(response);
879                cx.notify();
880            })
881        })
882    }
883
884    async fn fetch_models_request(
885        http_client: &HttpClientWithUrl,
886        token_provider: &TP,
887        auth_context: TP::AuthContext,
888    ) -> Result<ListModelsResponse> {
889        let url = http_client.build_zed_llm_url("/models", &[])?;
890        let mut response =
891            authenticated_llm_request(http_client, token_provider, auth_context, |token| {
892                Ok(http_client::Request::builder()
893                    .method(Method::GET)
894                    .header(CLIENT_SUPPORTS_X_AI_HEADER_NAME, "true")
895                    .uri(url.as_ref())
896                    .header("Authorization", format!("Bearer {token}"))
897                    .body(AsyncBody::empty())?)
898            })
899            .await
900            .context("failed to send list models request")?;
901
902        if response.status().is_success() {
903            let mut body = String::new();
904            response.body_mut().read_to_string(&mut body).await?;
905            Ok(serde_json::from_str(&body)?)
906        } else {
907            let mut body = String::new();
908            response.body_mut().read_to_string(&mut body).await?;
909            anyhow::bail!(
910                "error listing models.\nStatus: {:?}\nBody: {body}",
911                response.status(),
912            );
913        }
914    }
915
916    pub fn update_models(&mut self, response: ListModelsResponse) {
917        let models: Vec<_> = response.models.into_iter().map(Arc::new).collect();
918
919        self.default_model = models
920            .iter()
921            .find(|model| {
922                response
923                    .default_model
924                    .as_ref()
925                    .is_some_and(|default_model_id| &model.id == default_model_id)
926            })
927            .cloned();
928        self.default_fast_model = models
929            .iter()
930            .find(|model| {
931                response
932                    .default_fast_model
933                    .as_ref()
934                    .is_some_and(|default_fast_model_id| &model.id == default_fast_model_id)
935            })
936            .cloned();
937        self.recommended_models = response
938            .recommended_models
939            .iter()
940            .filter_map(|id| models.iter().find(|model| &model.id == id))
941            .cloned()
942            .collect();
943        self.models = models;
944    }
945
946    pub fn clear_models(&mut self) {
947        self.models.clear();
948        self.default_model = None;
949        self.default_fast_model = None;
950        self.recommended_models.clear();
951    }
952
953    pub fn create_model(
954        &self,
955        model: &Arc<cloud_llm_client::LanguageModel>,
956    ) -> Arc<dyn LanguageModel> {
957        Arc::new(CloudLanguageModel::<TP> {
958            id: LanguageModelId::from(model.id.0.to_string()),
959            model: model.clone(),
960            token_provider: self.token_provider.clone(),
961            http_client: self.http_client.clone(),
962            app_version: self.app_version.clone(),
963            request_limiter: RateLimiter::new(4),
964        })
965    }
966
967    pub fn models(&self) -> &[Arc<cloud_llm_client::LanguageModel>] {
968        &self.models
969    }
970
971    pub fn default_model(&self) -> Option<&Arc<cloud_llm_client::LanguageModel>> {
972        self.default_model.as_ref()
973    }
974
975    pub fn default_fast_model(&self) -> Option<&Arc<cloud_llm_client::LanguageModel>> {
976        self.default_fast_model.as_ref()
977    }
978
979    pub fn recommended_models(&self) -> &[Arc<cloud_llm_client::LanguageModel>] {
980        &self.recommended_models
981    }
982}
983
984pub fn map_cloud_completion_events<T, F>(
985    stream: Pin<Box<dyn Stream<Item = Result<CompletionEvent<T>, ResponseStreamError>> + Send>>,
986    provider: &LanguageModelProviderName,
987    mut map_callback: F,
988) -> BoxStream<'static, Result<LanguageModelCompletionEvent, LanguageModelCompletionError>>
989where
990    T: DeserializeOwned + 'static,
991    F: FnMut(T) -> Vec<Result<LanguageModelCompletionEvent, LanguageModelCompletionError>>
992        + Send
993        + 'static,
994{
995    let provider = provider.clone();
996    let mut stream = stream.fuse();
997
998    let mut saw_stream_ended = false;
999
1000    let mut done = false;
1001    let mut pending = VecDeque::new();
1002
1003    stream::poll_fn(move |cx| {
1004        loop {
1005            if let Some(item) = pending.pop_front() {
1006                return Poll::Ready(Some(item));
1007            }
1008
1009            if done {
1010                return Poll::Ready(None);
1011            }
1012
1013            match stream.poll_next_unpin(cx) {
1014                Poll::Ready(Some(event)) => {
1015                    let items = match event {
1016                        Err(error) => {
1017                            vec![Err(error.into_completion_error(provider.clone()))]
1018                        }
1019                        Ok(CompletionEvent::Status(CompletionRequestStatus::StreamEnded)) => {
1020                            saw_stream_ended = true;
1021                            vec![]
1022                        }
1023                        Ok(CompletionEvent::Status(status)) => {
1024                            LanguageModelCompletionEvent::from_completion_request_status(
1025                                status,
1026                                provider.clone(),
1027                            )
1028                            .transpose()
1029                            .map(|event| vec![event])
1030                            .unwrap_or_default()
1031                        }
1032                        Ok(CompletionEvent::Event(event)) => map_callback(event),
1033                    };
1034                    pending.extend(items);
1035                }
1036                Poll::Ready(None) => {
1037                    done = true;
1038
1039                    if !saw_stream_ended {
1040                        return Poll::Ready(Some(Err(
1041                            LanguageModelCompletionError::StreamEndedUnexpectedly {
1042                                provider: provider.clone(),
1043                            },
1044                        )));
1045                    }
1046                }
1047                Poll::Pending => return Poll::Pending,
1048            }
1049        }
1050    })
1051    .boxed()
1052}
1053
1054pub fn provider_name(
1055    provider: &cloud_llm_client::LanguageModelProvider,
1056) -> LanguageModelProviderName {
1057    match provider {
1058        cloud_llm_client::LanguageModelProvider::Anthropic => ANTHROPIC_PROVIDER_NAME,
1059        cloud_llm_client::LanguageModelProvider::OpenAi => OPEN_AI_PROVIDER_NAME,
1060        cloud_llm_client::LanguageModelProvider::Google => GOOGLE_PROVIDER_NAME,
1061        cloud_llm_client::LanguageModelProvider::XAi => X_AI_PROVIDER_NAME,
1062    }
1063}
1064
1065/// A failure while reading the streamed completion response body.
1066///
1067/// Kept as a typed error (rather than `anyhow::Error`) so the consumer can
1068/// attach the provider name and build a structured
1069/// [`LanguageModelCompletionError`] without a runtime downcast.
1070pub enum ResponseStreamError {
1071    Read(std::io::Error),
1072    Deserialize(serde_json::Error),
1073}
1074
1075impl ResponseStreamError {
1076    fn into_completion_error(
1077        self,
1078        provider: LanguageModelProviderName,
1079    ) -> LanguageModelCompletionError {
1080        match self {
1081            ResponseStreamError::Read(error) => {
1082                LanguageModelCompletionError::ApiReadResponseError { provider, error }
1083            }
1084            ResponseStreamError::Deserialize(error) => {
1085                LanguageModelCompletionError::DeserializeResponse { provider, error }
1086            }
1087        }
1088    }
1089}
1090
1091pub fn response_lines<T: DeserializeOwned>(
1092    response: Response<AsyncBody>,
1093    includes_status_messages: bool,
1094) -> impl Stream<Item = Result<CompletionEvent<T>, ResponseStreamError>> {
1095    futures::stream::try_unfold(
1096        (String::new(), BufReader::new(response.into_body())),
1097        move |(mut line, mut body)| async move {
1098            match body.read_line(&mut line).await {
1099                Ok(0) => Ok(None),
1100                Ok(_) => {
1101                    let event = if includes_status_messages {
1102                        serde_json::from_str::<CompletionEvent<T>>(&line)
1103                            .map_err(ResponseStreamError::Deserialize)?
1104                    } else {
1105                        CompletionEvent::Event(
1106                            serde_json::from_str::<T>(&line)
1107                                .map_err(ResponseStreamError::Deserialize)?,
1108                        )
1109                    };
1110
1111                    line.clear();
1112                    Ok(Some((event, (line, body))))
1113                }
1114                Err(error) => Err(ResponseStreamError::Read(error)),
1115            }
1116        },
1117    )
1118}
1119
1120#[cfg(test)]
1121mod tests {
1122    use super::*;
1123    use http_client::FakeHttpClient;
1124    use http_client::http::{HeaderMap, StatusCode};
1125    use language_model::{
1126        LanguageModelCompletionError, LanguageModelRequestMessage, MessageContent, Role, Speed,
1127    };
1128    use serde_json::json;
1129    use std::sync::Mutex;
1130
1131    #[gpui::test]
1132    async fn cloud_explicit_compaction_forwards_supported_request_fields(
1133        cx: &mut gpui::TestAppContext,
1134    ) {
1135        let captured_request = Arc::new(Mutex::new(None));
1136        let captured_request_for_handler = captured_request.clone();
1137        let http_client = FakeHttpClient::create(move |request| {
1138            let captured_request = captured_request_for_handler.clone();
1139            async move {
1140                let method = request.method().clone();
1141                let uri = request.uri().to_string();
1142                let authorization = request
1143                    .headers()
1144                    .get("Authorization")
1145                    .and_then(|value| value.to_str().ok())
1146                    .map(str::to_string);
1147                let requested_status_messages = request
1148                    .headers()
1149                    .contains_key(CLIENT_SUPPORTS_STATUS_MESSAGES_HEADER_NAME);
1150                let requested_stream_end = request
1151                    .headers()
1152                    .contains_key(CLIENT_SUPPORTS_STATUS_STREAM_ENDED_HEADER_NAME);
1153                let mut body = request.into_body();
1154                let mut body_text = String::new();
1155                body.read_to_string(&mut body_text).await?;
1156                *captured_request.lock().unwrap() = Some((
1157                    method,
1158                    uri,
1159                    authorization,
1160                    requested_status_messages,
1161                    requested_stream_end,
1162                    body_text,
1163                ));
1164
1165                Ok(http_client::Response::builder()
1166                    .status(200)
1167                    .body(AsyncBody::from(format!(
1168                        "{}\n",
1169                        json!({
1170                            "id": "resp_compact",
1171                            "created_at": 1_700_000_000,
1172                            "object": "response.compaction",
1173                            "output": [{
1174                                "type": "compaction",
1175                                "id": "cmp_manual",
1176                                "encrypted_content": "opaque-state"
1177                            }],
1178                            "usage": {
1179                                "input_tokens": 100,
1180                                "input_tokens_details": {"cached_tokens": 20},
1181                                "output_tokens": 10,
1182                                "output_tokens_details": {"reasoning_tokens": 5},
1183                                "total_tokens": 110
1184                            }
1185                        })
1186                    )))?)
1187            }
1188        });
1189        let model = cloud_test_model(http_client);
1190        let request = compact_test_request();
1191
1192        let result = model.compact(request, &cx.to_async()).await.unwrap();
1193
1194        assert_eq!(
1195            result.usage,
1196            language_model::TokenUsage {
1197                input_tokens: 80,
1198                output_tokens: 10,
1199                cache_creation_input_tokens: 0,
1200                cache_read_input_tokens: 20,
1201            }
1202        );
1203        let language_model::CompactedContext::ProviderState(state) = result.context else {
1204            panic!("expected provider compaction state");
1205        };
1206        assert_eq!(
1207            open_ai::responses::provider_compaction_items(&state, &OPEN_AI_PROVIDER_ID).unwrap(),
1208            Some(vec![json!({
1209                "type": "compaction",
1210                "id": "cmp_manual",
1211                "encrypted_content": "opaque-state"
1212            })])
1213        );
1214        let (method, uri, authorization, requested_status_messages, requested_stream_end, body) =
1215            captured_request.lock().unwrap().take().unwrap();
1216        assert_eq!(method, Method::POST);
1217        assert_eq!(uri, "http://test.example/completions/compact?");
1218        assert_eq!(authorization.as_deref(), Some("Bearer test-token"));
1219        assert!(!requested_status_messages);
1220        assert!(!requested_stream_end);
1221        let body = serde_json::from_str::<serde_json::Value>(&body).unwrap();
1222        assert_eq!(body["thread_id"], "thread-123");
1223        assert_eq!(body["provider"], "open_ai");
1224        assert_eq!(body["model"], "gpt-5.4");
1225        assert_eq!(
1226            body["provider_request"],
1227            json!({
1228                "model": "gpt-5.4",
1229                "input": [{
1230                    "type": "message",
1231                    "role": "user",
1232                    "content": [{
1233                        "type": "input_text",
1234                        "text": "Retain this context."
1235                    }]
1236                }],
1237                "prompt_cache_key": "thread-123",
1238                "service_tier": "priority"
1239            })
1240        );
1241    }
1242
1243    #[gpui::test]
1244    async fn cloud_explicit_compaction_rejects_output_without_compaction_item(
1245        cx: &mut gpui::TestAppContext,
1246    ) {
1247        let http_client = FakeHttpClient::create(|_| async move {
1248            Ok(http_client::Response::builder()
1249                .status(200)
1250                .body(AsyncBody::from(format!(
1251                    "{}\n",
1252                    json!({
1253                        "id": "resp_compact",
1254                        "created_at": 1_700_000_000,
1255                        "object": "response.compaction",
1256                        "output": [{
1257                            "type": "message",
1258                            "role": "assistant",
1259                            "content": "This is not an opaque compaction item."
1260                        }],
1261                        "usage": {
1262                            "input_tokens": 100,
1263                            "input_tokens_details": {"cached_tokens": 20},
1264                            "output_tokens": 10,
1265                            "output_tokens_details": {"reasoning_tokens": 5},
1266                            "total_tokens": 110
1267                        }
1268                    })
1269                )))?)
1270        });
1271        let model = cloud_test_model(http_client);
1272
1273        let error = model
1274            .compact(compact_test_request(), &cx.to_async())
1275            .await
1276            .unwrap_err();
1277
1278        assert!(
1279            matches!(&error, LanguageModelCompletionError::Other(_)),
1280            "expected invalid canonical output to be rejected, got {error:?}"
1281        );
1282        assert!(error.to_string().contains("compaction item"));
1283    }
1284
1285    #[test]
1286    fn test_api_error_conversion_with_upstream_http_error() {
1287        // upstream_http_error with 503 status should become ServerOverloaded
1288        let error_body = r#"{"code":"upstream_http_error","message":"Received an error from the Anthropic API: upstream connect error or disconnect/reset before headers, reset reason: connection timeout","upstream_status":503}"#;
1289
1290        let api_error = ApiError {
1291            status: StatusCode::INTERNAL_SERVER_ERROR,
1292            body: error_body.to_string(),
1293            headers: HeaderMap::new(),
1294        };
1295
1296        let completion_error: LanguageModelCompletionError = api_error.into();
1297
1298        match completion_error {
1299            LanguageModelCompletionError::UpstreamProviderError { message, .. } => {
1300                assert_eq!(
1301                    message,
1302                    "Received an error from the Anthropic API: upstream connect error or disconnect/reset before headers, reset reason: connection timeout"
1303                );
1304            }
1305            _ => panic!(
1306                "Expected UpstreamProviderError for upstream 503, got: {:?}",
1307                completion_error
1308            ),
1309        }
1310
1311        // upstream_http_error with 500 status should become ApiInternalServerError
1312        let error_body = r#"{"code":"upstream_http_error","message":"Received an error from the OpenAI API: internal server error","upstream_status":500}"#;
1313
1314        let api_error = ApiError {
1315            status: StatusCode::INTERNAL_SERVER_ERROR,
1316            body: error_body.to_string(),
1317            headers: HeaderMap::new(),
1318        };
1319
1320        let completion_error: LanguageModelCompletionError = api_error.into();
1321
1322        match completion_error {
1323            LanguageModelCompletionError::UpstreamProviderError { message, .. } => {
1324                assert_eq!(
1325                    message,
1326                    "Received an error from the OpenAI API: internal server error"
1327                );
1328            }
1329            _ => panic!(
1330                "Expected UpstreamProviderError for upstream 500, got: {:?}",
1331                completion_error
1332            ),
1333        }
1334
1335        // upstream_http_error with 429 status should become RateLimitExceeded
1336        let error_body = r#"{"code":"upstream_http_error","message":"Received an error from the Google API: rate limit exceeded","upstream_status":429}"#;
1337
1338        let api_error = ApiError {
1339            status: StatusCode::INTERNAL_SERVER_ERROR,
1340            body: error_body.to_string(),
1341            headers: HeaderMap::new(),
1342        };
1343
1344        let completion_error: LanguageModelCompletionError = api_error.into();
1345
1346        match completion_error {
1347            LanguageModelCompletionError::UpstreamProviderError { message, .. } => {
1348                assert_eq!(
1349                    message,
1350                    "Received an error from the Google API: rate limit exceeded"
1351                );
1352            }
1353            _ => panic!(
1354                "Expected UpstreamProviderError for upstream 429, got: {:?}",
1355                completion_error
1356            ),
1357        }
1358
1359        // Regular 500 error without upstream_http_error should remain ApiInternalServerError for Zed
1360        let error_body = "Regular internal server error";
1361
1362        let api_error = ApiError {
1363            status: StatusCode::INTERNAL_SERVER_ERROR,
1364            body: error_body.to_string(),
1365            headers: HeaderMap::new(),
1366        };
1367
1368        let completion_error: LanguageModelCompletionError = api_error.into();
1369
1370        match completion_error {
1371            LanguageModelCompletionError::ApiInternalServerError { provider, message } => {
1372                assert_eq!(provider, PROVIDER_NAME);
1373                assert_eq!(message, "Regular internal server error");
1374            }
1375            _ => panic!(
1376                "Expected ApiInternalServerError for regular 500, got: {:?}",
1377                completion_error
1378            ),
1379        }
1380
1381        // upstream_http_429 format should be converted to UpstreamProviderError
1382        let error_body = r#"{"code":"upstream_http_429","message":"Upstream Anthropic rate limit exceeded.","retry_after":30.5}"#;
1383
1384        let api_error = ApiError {
1385            status: StatusCode::INTERNAL_SERVER_ERROR,
1386            body: error_body.to_string(),
1387            headers: HeaderMap::new(),
1388        };
1389
1390        let completion_error: LanguageModelCompletionError = api_error.into();
1391
1392        match completion_error {
1393            LanguageModelCompletionError::UpstreamProviderError {
1394                message,
1395                status,
1396                retry_after,
1397            } => {
1398                assert_eq!(message, "Upstream Anthropic rate limit exceeded.");
1399                assert_eq!(status, StatusCode::TOO_MANY_REQUESTS);
1400                assert_eq!(retry_after, Some(Duration::from_secs_f64(30.5)));
1401            }
1402            _ => panic!(
1403                "Expected UpstreamProviderError for upstream_http_429, got: {:?}",
1404                completion_error
1405            ),
1406        }
1407
1408        // Invalid JSON in error body should fall back to regular error handling
1409        let error_body = "Not JSON at all";
1410
1411        let api_error = ApiError {
1412            status: StatusCode::INTERNAL_SERVER_ERROR,
1413            body: error_body.to_string(),
1414            headers: HeaderMap::new(),
1415        };
1416
1417        let completion_error: LanguageModelCompletionError = api_error.into();
1418
1419        match completion_error {
1420            LanguageModelCompletionError::ApiInternalServerError { provider, .. } => {
1421                assert_eq!(provider, PROVIDER_NAME);
1422            }
1423            _ => panic!(
1424                "Expected ApiInternalServerError for invalid JSON, got: {:?}",
1425                completion_error
1426            ),
1427        }
1428    }
1429
1430    #[test]
1431    fn test_response_stream_error_maps_to_structured_variant() {
1432        // Read/deserialize failures mid-stream must keep their structured
1433        // variant rather than collapsing into `Other` (the source of the
1434        // generic "Request failed." message).
1435        let read = ResponseStreamError::Read(std::io::Error::from(std::io::ErrorKind::BrokenPipe))
1436            .into_completion_error(PROVIDER_NAME);
1437        assert!(
1438            matches!(
1439                read,
1440                LanguageModelCompletionError::ApiReadResponseError { .. }
1441            ),
1442            "Expected ApiReadResponseError, got: {read:?}"
1443        );
1444
1445        let deserialize = ResponseStreamError::Deserialize(
1446            serde_json::from_str::<serde_json::Value>("not json").unwrap_err(),
1447        )
1448        .into_completion_error(PROVIDER_NAME);
1449        assert!(
1450            matches!(
1451                deserialize,
1452                LanguageModelCompletionError::DeserializeResponse { .. }
1453            ),
1454            "Expected DeserializeResponse, got: {deserialize:?}"
1455        );
1456    }
1457
1458    fn compact_test_request() -> LanguageModelRequest {
1459        LanguageModelRequest {
1460            thread_id: Some("thread-123".to_string()),
1461            messages: vec![LanguageModelRequestMessage {
1462                role: Role::User,
1463                content: vec![MessageContent::Text("Retain this context.".to_string())],
1464                cache: false,
1465                reasoning_details: None,
1466            }],
1467            speed: Some(Speed::Fast),
1468            ..Default::default()
1469        }
1470    }
1471
1472    fn cloud_test_model(
1473        http_client: Arc<HttpClientWithUrl>,
1474    ) -> CloudLanguageModel<TestTokenProvider> {
1475        CloudLanguageModel {
1476            id: LanguageModelId::from("gpt-5.4".to_string()),
1477            model: Arc::new(cloud_llm_client::LanguageModel {
1478                provider: cloud_llm_client::LanguageModelProvider::OpenAi,
1479                id: cloud_llm_client::LanguageModelId(Arc::from("gpt-5.4")),
1480                display_name: "GPT-5.4".to_string(),
1481                is_latest: true,
1482                max_token_count: 1_000_000,
1483                max_token_count_in_max_mode: None,
1484                max_output_tokens: 128_000,
1485                supports_tools: true,
1486                supports_images: true,
1487                supports_thinking: true,
1488                supports_disabling_thinking: true,
1489                supports_fast_mode: true,
1490                supports_server_side_compaction: true,
1491                supported_effort_levels: Vec::new(),
1492                supports_streaming_tools: true,
1493                supports_parallel_tool_calls: true,
1494                is_disabled: false,
1495                disabled_reason: None,
1496            }),
1497            token_provider: Arc::new(TestTokenProvider),
1498            http_client,
1499            app_version: None,
1500            request_limiter: RateLimiter::new(4),
1501        }
1502    }
1503
1504    struct TestTokenProvider;
1505
1506    impl CloudLlmTokenProvider for TestTokenProvider {
1507        type AuthContext = ();
1508
1509        fn auth_context(&self, _cx: &impl AppContext) -> Self::AuthContext {}
1510
1511        fn cached_token(
1512            &self,
1513            _auth_context: Self::AuthContext,
1514        ) -> BoxFuture<'static, Result<String>> {
1515            async { Ok("test-token".to_string()) }.boxed()
1516        }
1517
1518        fn refresh_token(
1519            &self,
1520            _auth_context: Self::AuthContext,
1521        ) -> BoxFuture<'static, Result<String>> {
1522            async { Ok("refreshed-test-token".to_string()) }.boxed()
1523        }
1524
1525        fn has_data_retention_consent(&self, _cx: &impl AppContext) -> bool {
1526            false
1527        }
1528    }
1529}
1530
Served at tenant.openagents/omega Member data and write actions are omitted.