Skip to repository content

tenant.openagents/omega

No repository description is available.

OpenAgents Git authority 2026-07-28T02:56:22.856Z 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_model_core.rs

927 lines · 30.8 KB · rust
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
Served at tenant.openagents/omega Member data and write actions are omitted.