Skip to repository content1530 lines · 58.6 KB · rust
tenant.openagents/omega
No repository description is available.
OpenAgents Git authority 2026-07-28T03:30:34.297Z Public web read
NIP-34 coordinate
30617:7649603503856e5148d571eac2766b288a8ff1e9e35d380337a1d2b0015b4f92:omegaMaintainersHidden in public view
References2 branches · 1 tag
Read-only clone
git clone https://openagents.com/git/tenant.openagents/omega.gitBrowse files
language_models_cloud.rs
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