Skip to repository content927 lines · 30.8 KB · rust
tenant.openagents/omega
No repository description is available.
OpenAgents Git authority 2026-07-28T04:01:43.311Z 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_model_core.rs
1mod provider;
2mod rate_limiter;
3mod request;
4mod role;
5pub mod tool_schema;
6pub mod util;
7
8use anyhow::{Context as _, Result, anyhow};
9use cloud_llm_client::CompletionRequestStatus;
10use http_client::{StatusCode, http};
11use schemars::JsonSchema;
12use serde::{Deserialize, Serialize};
13use std::ops::{Add, Sub};
14use std::str::FromStr;
15use std::sync::Arc;
16use std::time::Duration;
17use std::{fmt, io};
18use thiserror::Error;
19fn is_default<T: Default + PartialEq>(value: &T) -> bool {
20 *value == T::default()
21}
22
23pub use crate::provider::*;
24pub use crate::rate_limiter::*;
25pub use crate::request::*;
26pub use crate::role::*;
27pub use crate::tool_schema::LanguageModelToolSchemaFormat;
28pub use crate::util::{
29 fix_streamed_json, is_context_window_exceeded_message, parse_prompt_too_long,
30 parse_tool_arguments,
31};
32pub use gpui_shared_string::SharedString;
33
34/// A completion event from a language model.
35#[derive(Debug, PartialEq, Clone, Serialize, Deserialize)]
36pub enum LanguageModelCompletionEvent {
37 Queued {
38 position: usize,
39 },
40 Started,
41 Stop(StopReason),
42 Text(String),
43 Thinking {
44 text: String,
45 signature: Option<String>,
46 },
47 RedactedThinking {
48 data: String,
49 },
50 ToolUse(LanguageModelToolUse),
51 ToolUseJsonParseError {
52 id: LanguageModelToolUseId,
53 tool_name: Arc<str>,
54 raw_input: Arc<str>,
55 json_parse_error: String,
56 },
57 StartMessage {
58 message_id: String,
59 },
60 ReasoningDetails(serde_json::Value),
61 UsageUpdate(TokenUsage),
62 Compaction(CompactionUpdate),
63}
64
65#[derive(Debug, PartialEq, Clone, Serialize, Deserialize)]
66pub enum CompactionUpdate {
67 /// A streamed response has started producing replacement context.
68 Started,
69 /// A chunk of a natural-language summary, suitable for incremental display.
70 SummaryDelta(Arc<str>),
71 /// The complete context to persist and use in subsequent requests.
72 Finished(CompactedContext),
73 /// The provider abandoned the compaction without producing replacement
74 /// context. This is a documented outcome, not a protocol error: the
75 /// conversation simply continues on the uncompacted transcript.
76 Failed,
77}
78
79impl LanguageModelCompletionEvent {
80 pub fn from_completion_request_status(
81 status: CompletionRequestStatus,
82 upstream_provider: LanguageModelProviderName,
83 ) -> Result<Option<Self>, LanguageModelCompletionError> {
84 match status {
85 CompletionRequestStatus::Queued { position } => {
86 Ok(Some(LanguageModelCompletionEvent::Queued { position }))
87 }
88 CompletionRequestStatus::Started => Ok(Some(LanguageModelCompletionEvent::Started)),
89 CompletionRequestStatus::Unknown | CompletionRequestStatus::StreamEnded => Ok(None),
90 CompletionRequestStatus::Failed {
91 code,
92 message,
93 request_id: _,
94 retry_after,
95 } => Err(LanguageModelCompletionError::from_cloud_failure(
96 upstream_provider,
97 code,
98 message,
99 retry_after.map(Duration::from_secs_f64),
100 )),
101 }
102 }
103}
104
105#[derive(Error, Debug)]
106pub enum LanguageModelCompletionError {
107 #[error("prompt too large for context window")]
108 PromptTooLarge { tokens: Option<u64> },
109 /// The model requires the user to consent to the upstream provider
110 /// retaining inference logs (see `LanguageModel::requires_data_retention`)
111 /// and that consent has not been given.
112 #[error(
113 "{model_name} cannot be offered with Zero Data Retention. \
114 Anthropic will retain inference logs."
115 )]
116 DataRetentionConsentRequired { model_name: String },
117 #[error("missing {provider} API key")]
118 NoApiKey { provider: LanguageModelProviderName },
119 #[error("{provider}'s API rate limit exceeded")]
120 RateLimitExceeded {
121 provider: LanguageModelProviderName,
122 retry_after: Option<Duration>,
123 },
124 #[error("{provider}'s API servers are overloaded right now")]
125 ServerOverloaded {
126 provider: LanguageModelProviderName,
127 retry_after: Option<Duration>,
128 },
129 #[error("{provider}'s API server reported an internal server error: {message}")]
130 ApiInternalServerError {
131 provider: LanguageModelProviderName,
132 message: String,
133 },
134 #[error("{message}")]
135 UpstreamProviderError {
136 message: String,
137 status: StatusCode,
138 retry_after: Option<Duration>,
139 },
140 #[error("HTTP response error from {provider}'s API: status {status_code} - {message:?}")]
141 HttpResponseError {
142 provider: LanguageModelProviderName,
143 status_code: StatusCode,
144 message: String,
145 },
146 #[error("invalid request format to {provider}'s API: {message}")]
147 BadRequestFormat {
148 provider: LanguageModelProviderName,
149 message: String,
150 },
151 #[error("authentication error with {provider}'s API: {message}")]
152 AuthenticationError {
153 provider: LanguageModelProviderName,
154 message: String,
155 },
156 #[error("Permission error with {provider}'s API: {message}")]
157 PermissionError {
158 provider: LanguageModelProviderName,
159 message: String,
160 },
161 #[error("language model provider API endpoint not found")]
162 ApiEndpointNotFound { provider: LanguageModelProviderName },
163 #[error("I/O error reading response from {provider}'s API")]
164 ApiReadResponseError {
165 provider: LanguageModelProviderName,
166 #[source]
167 error: io::Error,
168 },
169 #[error("error serializing request to {provider} API")]
170 SerializeRequest {
171 provider: LanguageModelProviderName,
172 #[source]
173 error: serde_json::Error,
174 },
175 #[error("error building request body to {provider} API")]
176 BuildRequestBody {
177 provider: LanguageModelProviderName,
178 #[source]
179 error: http::Error,
180 },
181 #[error("error sending HTTP request to {provider} API")]
182 HttpSend {
183 provider: LanguageModelProviderName,
184 #[source]
185 error: anyhow::Error,
186 },
187 #[error("error deserializing {provider} API response")]
188 DeserializeResponse {
189 provider: LanguageModelProviderName,
190 #[source]
191 error: serde_json::Error,
192 },
193 #[error("stream from {provider} ended unexpectedly")]
194 StreamEndedUnexpectedly { provider: LanguageModelProviderName },
195 #[error("payment required to use this language model; please upgrade your account")]
196 PaymentRequired,
197 #[error(transparent)]
198 Other(#[from] anyhow::Error),
199}
200
201impl LanguageModelCompletionError {
202 fn parse_upstream_error_json(message: &str) -> Option<(StatusCode, String)> {
203 let error_json = serde_json::from_str::<serde_json::Value>(message).ok()?;
204 let upstream_status = error_json
205 .get("upstream_status")
206 .and_then(|v| v.as_u64())
207 .and_then(|status| u16::try_from(status).ok())
208 .and_then(|status| StatusCode::from_u16(status).ok())?;
209 let inner_message = error_json
210 .get("message")
211 .and_then(|v| v.as_str())
212 .unwrap_or(message)
213 .to_string();
214 Some((upstream_status, inner_message))
215 }
216
217 pub fn from_cloud_failure(
218 upstream_provider: LanguageModelProviderName,
219 code: String,
220 message: String,
221 retry_after: Option<Duration>,
222 ) -> Self {
223 if let Some(tokens) = parse_prompt_too_long(&message) {
224 Self::PromptTooLarge {
225 tokens: Some(tokens),
226 }
227 } else if code == "upstream_http_error" {
228 if let Some((upstream_status, inner_message)) =
229 Self::parse_upstream_error_json(&message)
230 {
231 return Self::from_http_status(
232 upstream_provider,
233 upstream_status,
234 inner_message,
235 retry_after,
236 );
237 }
238 anyhow!("completion request failed, code: {code}, message: {message}").into()
239 } else if let Some(status_code) = code
240 .strip_prefix("upstream_http_")
241 .and_then(|code| StatusCode::from_str(code).ok())
242 {
243 Self::from_http_status(upstream_provider, status_code, message, retry_after)
244 } else if let Some(status_code) = code
245 .strip_prefix("http_")
246 .and_then(|code| StatusCode::from_str(code).ok())
247 {
248 Self::from_http_status(ZED_CLOUD_PROVIDER_NAME, status_code, message, retry_after)
249 } else {
250 anyhow!("completion request failed, code: {code}, message: {message}").into()
251 }
252 }
253
254 pub fn from_http_status(
255 provider: LanguageModelProviderName,
256 status_code: StatusCode,
257 message: String,
258 retry_after: Option<Duration>,
259 ) -> Self {
260 match status_code {
261 StatusCode::BAD_REQUEST => {
262 if is_context_window_exceeded_message(&message) {
263 Self::PromptTooLarge { tokens: None }
264 } else {
265 Self::BadRequestFormat { provider, message }
266 }
267 }
268 StatusCode::UNAUTHORIZED => Self::AuthenticationError { provider, message },
269 StatusCode::FORBIDDEN => Self::PermissionError { provider, message },
270 StatusCode::NOT_FOUND => Self::ApiEndpointNotFound { provider },
271 StatusCode::PAYLOAD_TOO_LARGE => Self::PromptTooLarge {
272 tokens: parse_prompt_too_long(&message),
273 },
274 StatusCode::TOO_MANY_REQUESTS => Self::RateLimitExceeded {
275 provider,
276 retry_after,
277 },
278 StatusCode::INTERNAL_SERVER_ERROR => Self::ApiInternalServerError { provider, message },
279 StatusCode::SERVICE_UNAVAILABLE => Self::ServerOverloaded {
280 provider,
281 retry_after,
282 },
283 _ if status_code.as_u16() == 529 => Self::ServerOverloaded {
284 provider,
285 retry_after,
286 },
287 _ => Self::HttpResponseError {
288 provider,
289 status_code,
290 message,
291 },
292 }
293 }
294}
295
296#[derive(Debug, PartialEq, Clone, Copy, Serialize, Deserialize)]
297#[serde(rename_all = "snake_case")]
298pub enum StopReason {
299 EndTurn,
300 MaxTokens,
301 ToolUse,
302 Refusal,
303}
304
305#[derive(Debug, PartialEq, Clone, Copy, Serialize, Deserialize, Default)]
306pub struct TokenUsage {
307 #[serde(default, skip_serializing_if = "is_default")]
308 pub input_tokens: u64,
309 #[serde(default, skip_serializing_if = "is_default")]
310 pub output_tokens: u64,
311 #[serde(default, skip_serializing_if = "is_default")]
312 pub cache_creation_input_tokens: u64,
313 #[serde(default, skip_serializing_if = "is_default")]
314 pub cache_read_input_tokens: u64,
315}
316
317impl TokenUsage {
318 pub fn total_tokens(&self) -> u64 {
319 self.input_tokens
320 + self.output_tokens
321 + self.cache_read_input_tokens
322 + self.cache_creation_input_tokens
323 }
324}
325
326impl Add<TokenUsage> for TokenUsage {
327 type Output = Self;
328
329 fn add(self, other: Self) -> Self {
330 Self {
331 input_tokens: self.input_tokens + other.input_tokens,
332 output_tokens: self.output_tokens + other.output_tokens,
333 cache_creation_input_tokens: self.cache_creation_input_tokens
334 + other.cache_creation_input_tokens,
335 cache_read_input_tokens: self.cache_read_input_tokens + other.cache_read_input_tokens,
336 }
337 }
338}
339
340impl Sub<TokenUsage> for TokenUsage {
341 type Output = Self;
342
343 fn sub(self, other: Self) -> Self {
344 Self {
345 input_tokens: self.input_tokens - other.input_tokens,
346 output_tokens: self.output_tokens - other.output_tokens,
347 cache_creation_input_tokens: self.cache_creation_input_tokens
348 - other.cache_creation_input_tokens,
349 cache_read_input_tokens: self.cache_read_input_tokens - other.cache_read_input_tokens,
350 }
351 }
352}
353
354#[derive(Debug, PartialEq, Eq, Hash, Clone, Serialize, Deserialize)]
355pub struct LanguageModelToolUseId(Arc<str>);
356
357impl fmt::Display for LanguageModelToolUseId {
358 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
359 write!(f, "{}", self.0)
360 }
361}
362
363impl<T> From<T> for LanguageModelToolUseId
364where
365 T: Into<Arc<str>>,
366{
367 fn from(value: T) -> Self {
368 Self(value.into())
369 }
370}
371
372#[derive(Debug, PartialEq, Eq, Hash, Clone, Serialize, Deserialize)]
373pub struct LanguageModelToolUse {
374 pub id: LanguageModelToolUseId,
375 pub name: Arc<str>,
376 pub raw_input: String,
377 pub input: LanguageModelToolUseInput,
378 pub is_input_complete: bool,
379 /// Thought signature the model sent us. Some models require that this
380 /// signature be preserved and sent back in conversation history for validation.
381 pub thought_signature: Option<String>,
382}
383
384#[derive(Debug, PartialEq, Eq, Hash, Clone)]
385pub enum LanguageModelToolUseInput {
386 Json(serde_json::Value),
387 Text(String),
388}
389
390impl Serialize for LanguageModelToolUseInput {
391 fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
392 where
393 S: serde::Serializer,
394 {
395 use serde::ser::SerializeStruct;
396
397 let mut state = serializer.serialize_struct("LanguageModelToolUseInput", 2)?;
398 match self {
399 Self::Json(input) => {
400 state.serialize_field("type", "json")?;
401 state.serialize_field("value", input)?;
402 }
403 Self::Text(input) => {
404 state.serialize_field("type", "text")?;
405 state.serialize_field("value", input)?;
406 }
407 }
408 state.end()
409 }
410}
411
412impl<'de> Deserialize<'de> for LanguageModelToolUseInput {
413 fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
414 where
415 D: serde::Deserializer<'de>,
416 {
417 let value = serde_json::Value::deserialize(deserializer)?;
418 if let Some(object) = value.as_object()
419 && object.len() == 2
420 && let Some(input_type) = object.get("type").and_then(|value| value.as_str())
421 && let Some(input) = object.get("value")
422 {
423 return match input_type {
424 "json" => Ok(Self::Json(input.clone())),
425 "text" => input
426 .as_str()
427 .map(|input| Self::Text(input.to_string()))
428 .ok_or_else(|| serde::de::Error::custom("text tool input must be a string")),
429 _ => Ok(Self::Json(value)),
430 };
431 }
432
433 Ok(Self::Json(value))
434 }
435}
436
437impl LanguageModelToolUseInput {
438 pub fn as_json(&self) -> Option<&serde_json::Value> {
439 match self {
440 Self::Json(input) => Some(input),
441 Self::Text(_) => None,
442 }
443 }
444
445 /// Typed parsing for JSON tool inputs; freeform (Text) inputs always error.
446 ///
447 /// Callers wanting the raw value should use [`Self::as_json`] or [`Self::into_json`].
448 pub fn parse<T: serde::de::DeserializeOwned>(&self) -> Result<T> {
449 match self {
450 Self::Json(input) => {
451 serde_json::from_value(input.clone()).context("failed to parse JSON tool input")
452 }
453 Self::Text(_) => Err(anyhow!("custom tool text input cannot be parsed as JSON")),
454 }
455 }
456
457 pub fn into_json(self) -> Result<serde_json::Value> {
458 match self {
459 Self::Json(input) => Ok(input),
460 Self::Text(_) => Err(anyhow!("custom tool text input cannot be used as JSON")),
461 }
462 }
463
464 pub fn to_display_json(&self) -> serde_json::Value {
465 match self {
466 Self::Json(input) => input.clone(),
467 Self::Text(input) => serde_json::Value::String(input.clone()),
468 }
469 }
470}
471
472#[derive(Debug, Clone)]
473pub struct LanguageModelEffortLevel {
474 pub name: SharedString,
475 pub value: SharedString,
476 pub is_default: bool,
477}
478
479/// An error that occurred when trying to authenticate the language model provider.
480#[derive(Debug, Error)]
481pub enum AuthenticateError {
482 #[error("connection refused")]
483 ConnectionRefused,
484 #[error("credentials not found")]
485 CredentialsNotFound,
486 #[error(transparent)]
487 Other(#[from] anyhow::Error),
488}
489
490#[derive(Clone, Eq, PartialEq, Hash, Debug, Ord, PartialOrd, Serialize, Deserialize)]
491pub struct LanguageModelId(pub SharedString);
492
493#[derive(Clone, Eq, PartialEq, Hash, Debug, Ord, PartialOrd)]
494pub struct LanguageModelName(pub SharedString);
495
496#[derive(Clone, Eq, PartialEq, Hash, Debug, Ord, PartialOrd, Serialize, Deserialize)]
497pub struct LanguageModelProviderId(pub SharedString);
498
499#[derive(Clone, Eq, PartialEq, Hash, Debug, Ord, PartialOrd)]
500pub struct LanguageModelProviderName(pub SharedString);
501
502impl LanguageModelProviderId {
503 pub const fn new(id: &'static str) -> Self {
504 Self(SharedString::new_static(id))
505 }
506}
507
508impl LanguageModelProviderName {
509 pub const fn new(id: &'static str) -> Self {
510 Self(SharedString::new_static(id))
511 }
512}
513
514impl fmt::Display for LanguageModelProviderId {
515 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
516 write!(f, "{}", self.0)
517 }
518}
519
520impl fmt::Display for LanguageModelProviderName {
521 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
522 write!(f, "{}", self.0)
523 }
524}
525
526impl From<String> for LanguageModelId {
527 fn from(value: String) -> Self {
528 Self(SharedString::from(value))
529 }
530}
531
532impl From<String> for LanguageModelName {
533 fn from(value: String) -> Self {
534 Self(SharedString::from(value))
535 }
536}
537
538impl From<String> for LanguageModelProviderId {
539 fn from(value: String) -> Self {
540 Self(SharedString::from(value))
541 }
542}
543
544impl From<String> for LanguageModelProviderName {
545 fn from(value: String) -> Self {
546 Self(SharedString::from(value))
547 }
548}
549
550impl From<Arc<str>> for LanguageModelProviderId {
551 fn from(value: Arc<str>) -> Self {
552 Self(SharedString::from(value))
553 }
554}
555
556impl From<Arc<str>> for LanguageModelProviderName {
557 fn from(value: Arc<str>) -> Self {
558 Self(SharedString::from(value))
559 }
560}
561
562/// Settings-layer–free model mode enum.
563///
564/// Mirrors the shape of `settings_content::ModelMode` but lives here so that
565/// crates below the settings layer can reference it.
566#[derive(Copy, Clone, Debug, Default, PartialEq, Eq, Serialize, Deserialize, JsonSchema)]
567#[serde(tag = "type", rename_all = "lowercase")]
568pub enum ModelMode {
569 #[default]
570 Default,
571 Thinking {
572 budget_tokens: Option<u32>,
573 },
574 Adaptive,
575}
576
577/// Settings-layer–free reasoning-effort enum.
578///
579/// Mirrors the shape of `settings_content::OpenAiReasoningEffort` but lives
580/// here so that crates below the settings layer can reference it.
581#[derive(
582 Debug, Copy, Clone, PartialEq, Eq, Serialize, Deserialize, JsonSchema, strum::EnumString,
583)]
584#[serde(rename_all = "lowercase")]
585#[strum(serialize_all = "lowercase")]
586pub enum ReasoningEffort {
587 None,
588 Minimal,
589 Low,
590 Medium,
591 High,
592 XHigh,
593 Max,
594}
595
596impl ReasoningEffort {
597 pub const OPENAI_COMPATIBLE_SELECTABLE: [Self; 6] = [
598 Self::Minimal,
599 Self::Low,
600 Self::Medium,
601 Self::High,
602 Self::XHigh,
603 Self::Max,
604 ];
605
606 pub fn label(self) -> &'static str {
607 match self {
608 Self::None => "None",
609 Self::Minimal => "Minimal",
610 Self::Low => "Low",
611 Self::Medium => "Medium",
612 Self::High => "High",
613 Self::XHigh => "Extra High",
614 Self::Max => "Max",
615 }
616 }
617
618 pub fn value(self) -> &'static str {
619 match self {
620 Self::None => "none",
621 Self::Minimal => "minimal",
622 Self::Low => "low",
623 Self::Medium => "medium",
624 Self::High => "high",
625 Self::XHigh => "xhigh",
626 Self::Max => "max",
627 }
628 }
629}
630
631#[cfg(test)]
632mod tests {
633 use super::*;
634
635 #[test]
636 fn test_from_cloud_failure_with_upstream_http_error() {
637 let error = LanguageModelCompletionError::from_cloud_failure(
638 String::from("anthropic").into(),
639 "upstream_http_error".to_string(),
640 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}"#.to_string(),
641 None,
642 );
643
644 match error {
645 LanguageModelCompletionError::ServerOverloaded { provider, .. } => {
646 assert_eq!(provider.0, "anthropic");
647 }
648 _ => panic!(
649 "Expected ServerOverloaded error for 503 status, got: {:?}",
650 error
651 ),
652 }
653
654 let error = LanguageModelCompletionError::from_cloud_failure(
655 String::from("anthropic").into(),
656 "upstream_http_error".to_string(),
657 r#"{"code":"upstream_http_error","message":"Internal server error","upstream_status":500}"#.to_string(),
658 None,
659 );
660
661 match error {
662 LanguageModelCompletionError::ApiInternalServerError { provider, message } => {
663 assert_eq!(provider.0, "anthropic");
664 assert_eq!(message, "Internal server error");
665 }
666 _ => panic!(
667 "Expected ApiInternalServerError for 500 status, got: {:?}",
668 error
669 ),
670 }
671 }
672
673 #[test]
674 fn test_from_http_status_maps_context_length_exceeded_to_prompt_too_large() {
675 let error = LanguageModelCompletionError::from_http_status(
676 String::from("OpenAI").into(),
677 StatusCode::BAD_REQUEST,
678 r#"{"error":{"type":"invalid_request_error","code":"context_length_exceeded","message":"Your input exceeds the context window of this model. Please adjust your input and try again.","param":"input"}}"#.to_string(),
679 None,
680 );
681
682 assert!(matches!(
683 error,
684 LanguageModelCompletionError::PromptTooLarge { tokens: None }
685 ));
686
687 let error = LanguageModelCompletionError::from_http_status(
688 String::from("OpenAI").into(),
689 StatusCode::BAD_REQUEST,
690 "Invalid request.".to_string(),
691 None,
692 );
693
694 assert!(matches!(
695 error,
696 LanguageModelCompletionError::BadRequestFormat { .. }
697 ));
698 }
699
700 #[test]
701 fn test_from_cloud_failure_with_standard_format() {
702 let error = LanguageModelCompletionError::from_cloud_failure(
703 String::from("anthropic").into(),
704 "upstream_http_503".to_string(),
705 "Service unavailable".to_string(),
706 None,
707 );
708
709 match error {
710 LanguageModelCompletionError::ServerOverloaded { provider, .. } => {
711 assert_eq!(provider.0, "anthropic");
712 }
713 _ => panic!("Expected ServerOverloaded error for upstream_http_503"),
714 }
715 }
716
717 #[test]
718 fn test_upstream_http_error_connection_timeout() {
719 let error = LanguageModelCompletionError::from_cloud_failure(
720 String::from("anthropic").into(),
721 "upstream_http_error".to_string(),
722 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}"#.to_string(),
723 None,
724 );
725
726 match error {
727 LanguageModelCompletionError::ServerOverloaded { provider, .. } => {
728 assert_eq!(provider.0, "anthropic");
729 }
730 _ => panic!(
731 "Expected ServerOverloaded error for connection timeout with 503 status, got: {:?}",
732 error
733 ),
734 }
735
736 let error = LanguageModelCompletionError::from_cloud_failure(
737 String::from("anthropic").into(),
738 "upstream_http_error".to_string(),
739 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":500}"#.to_string(),
740 None,
741 );
742
743 match error {
744 LanguageModelCompletionError::ApiInternalServerError { provider, message } => {
745 assert_eq!(provider.0, "anthropic");
746 assert_eq!(
747 message,
748 "Received an error from the Anthropic API: upstream connect error or disconnect/reset before headers. reset reason: connection timeout"
749 );
750 }
751 _ => panic!(
752 "Expected ApiInternalServerError for connection timeout with 500 status, got: {:?}",
753 error
754 ),
755 }
756 }
757
758 #[test]
759 fn test_language_model_tool_use_serializes_with_signature() {
760 use serde_json::json;
761
762 let tool_use = LanguageModelToolUse {
763 id: LanguageModelToolUseId::from("test_id"),
764 name: "test_tool".into(),
765 raw_input: json!({"arg": "value"}).to_string(),
766 input: LanguageModelToolUseInput::Json(json!({"arg": "value"})),
767 is_input_complete: true,
768 thought_signature: Some("test_signature".to_string()),
769 };
770
771 let serialized = serde_json::to_value(&tool_use).unwrap();
772
773 assert_eq!(serialized["id"], "test_id");
774 assert_eq!(serialized["name"], "test_tool");
775 assert_eq!(serialized["thought_signature"], "test_signature");
776 }
777
778 #[test]
779 fn test_language_model_tool_use_deserializes_with_missing_signature() {
780 use serde_json::json;
781
782 let json = json!({
783 "id": "test_id",
784 "name": "test_tool",
785 "raw_input": "{\"arg\":\"value\"}",
786 "input": {"arg": "value"},
787 "is_input_complete": true
788 });
789
790 let tool_use: LanguageModelToolUse = serde_json::from_value(json).unwrap();
791
792 assert_eq!(tool_use.id, LanguageModelToolUseId::from("test_id"));
793 assert_eq!(tool_use.name.as_ref(), "test_tool");
794 assert_eq!(
795 tool_use.input,
796 LanguageModelToolUseInput::Json(json!({"arg": "value"}))
797 );
798 assert_eq!(tool_use.thought_signature, None);
799 }
800
801 #[test]
802 fn test_language_model_tool_use_input_round_trips_json() {
803 use serde_json::json;
804
805 let input = LanguageModelToolUseInput::Json(json!({"arg": "value"}));
806 let serialized = serde_json::to_value(&input).unwrap();
807 assert_eq!(
808 serialized,
809 json!({
810 "type": "json",
811 "value": {"arg": "value"}
812 })
813 );
814
815 let deserialized: LanguageModelToolUseInput = serde_json::from_value(serialized).unwrap();
816 assert_eq!(deserialized, input);
817 }
818
819 #[test]
820 fn test_language_model_tool_use_input_round_trips_text() {
821 use serde_json::json;
822
823 let input = LanguageModelToolUseInput::Text("raw custom input".to_string());
824 let serialized = serde_json::to_value(&input).unwrap();
825 assert_eq!(
826 serialized,
827 json!({
828 "type": "text",
829 "value": "raw custom input"
830 })
831 );
832
833 let deserialized: LanguageModelToolUseInput = serde_json::from_value(serialized).unwrap();
834 assert_eq!(deserialized, input);
835 }
836
837 #[test]
838 fn test_language_model_tool_use_input_parse() {
839 use serde_json::json;
840
841 #[derive(Debug, Deserialize, PartialEq)]
842 struct TestInput {
843 arg: String,
844 }
845
846 let parsed: TestInput = LanguageModelToolUseInput::Json(json!({"arg": "value"}))
847 .parse()
848 .unwrap();
849 assert_eq!(
850 parsed,
851 TestInput {
852 arg: "value".to_string()
853 }
854 );
855
856 let error = LanguageModelToolUseInput::Text("raw custom input".to_string())
857 .parse::<TestInput>()
858 .unwrap_err();
859 assert!(
860 error
861 .to_string()
862 .contains("custom tool text input cannot be parsed as JSON")
863 );
864 }
865
866 #[test]
867 fn test_language_model_tool_use_input_deserializes_legacy_plain_json_as_json() {
868 use serde_json::json;
869
870 let deserialized: LanguageModelToolUseInput =
871 serde_json::from_value(json!({"arg": "value"})).unwrap();
872 assert_eq!(
873 deserialized,
874 LanguageModelToolUseInput::Json(json!({"arg": "value"}))
875 );
876
877 let deserialized: LanguageModelToolUseInput =
878 serde_json::from_value(json!("legacy string argument")).unwrap();
879 assert_eq!(
880 deserialized,
881 LanguageModelToolUseInput::Json(json!("legacy string argument"))
882 );
883 }
884
885 #[test]
886 fn test_language_model_tool_use_round_trip_with_signature() {
887 use serde_json::json;
888
889 let original = LanguageModelToolUse {
890 id: LanguageModelToolUseId::from("round_trip_id"),
891 name: "round_trip_tool".into(),
892 raw_input: json!({"key": "value"}).to_string(),
893 input: LanguageModelToolUseInput::Json(json!({"key": "value"})),
894 is_input_complete: true,
895 thought_signature: Some("round_trip_sig".to_string()),
896 };
897
898 let serialized = serde_json::to_value(&original).unwrap();
899 let deserialized: LanguageModelToolUse = serde_json::from_value(serialized).unwrap();
900
901 assert_eq!(deserialized.id, original.id);
902 assert_eq!(deserialized.name, original.name);
903 assert_eq!(deserialized.thought_signature, original.thought_signature);
904 }
905
906 #[test]
907 fn test_language_model_tool_use_round_trip_without_signature() {
908 use serde_json::json;
909
910 let original = LanguageModelToolUse {
911 id: LanguageModelToolUseId::from("no_sig_id"),
912 name: "no_sig_tool".into(),
913 raw_input: json!({"arg": "value"}).to_string(),
914 input: LanguageModelToolUseInput::Json(json!({"arg": "value"})),
915 is_input_complete: true,
916 thought_signature: None,
917 };
918
919 let serialized = serde_json::to_value(&original).unwrap();
920 let deserialized: LanguageModelToolUse = serde_json::from_value(serialized).unwrap();
921
922 assert_eq!(deserialized.id, original.id);
923 assert_eq!(deserialized.name, original.name);
924 assert_eq!(deserialized.thought_signature, None);
925 }
926}
927