Skip to main content

rig_core/providers/openai/
embedding.rs

1use crate::json_utils;
2use serde::{Deserialize, Serialize};
3use std::fmt;
4
5/// `text-embedding-3-large` embedding model
6pub const TEXT_EMBEDDING_3_LARGE: &str = "text-embedding-3-large";
7/// `text-embedding-3-small` embedding model
8pub const TEXT_EMBEDDING_3_SMALL: &str = "text-embedding-3-small";
9/// `text-embedding-ada-002` embedding model
10pub const TEXT_EMBEDDING_ADA_002: &str = "text-embedding-ada-002";
11
12#[derive(Debug, Deserialize)]
13pub struct EmbeddingResponse {
14    pub object: String,
15    pub data: Vec<EmbeddingData>,
16    pub model: String,
17    pub usage: Usage,
18}
19
20/// Typed raw response from [`Embeddings`](super::wire::Embeddings).
21/// Missing object and model fields default to empty strings; usage is optional.
22#[derive(Debug, Clone, Serialize, Deserialize)]
23pub struct CompatibleEmbeddingResponse {
24    #[serde(default)]
25    pub object: String,
26    pub data: Vec<EmbeddingData>,
27    #[serde(default)]
28    pub model: String,
29    #[serde(default)]
30    pub usage: Option<Usage>,
31}
32
33#[derive(Debug, Deserialize, Clone, Copy, PartialEq, Eq, Serialize)]
34#[serde(rename_all = "snake_case")]
35pub enum EncodingFormat {
36    Float,
37    Base64,
38}
39
40/// One embedded input. Missing object and index fields default to empty and zero.
41#[derive(Debug, Clone, Serialize, Deserialize)]
42pub struct EmbeddingData {
43    #[serde(default)]
44    pub object: String,
45    pub embedding: Vec<serde_json::Number>,
46    #[serde(default)]
47    pub index: usize,
48}
49
50/// Return default dimensions for a known model identifier, or `None`.
51pub(crate) fn model_dimensions_from_identifier(identifier: &str) -> Option<usize> {
52    match identifier {
53        TEXT_EMBEDDING_3_LARGE => Some(3_072),
54        TEXT_EMBEDDING_3_SMALL | TEXT_EMBEDDING_ADA_002 => Some(1_536),
55        _ => None,
56    }
57}
58
59#[derive(Clone, Copy, Debug, Deserialize, Serialize, Default)]
60pub struct PromptTokensDetails {
61    /// Cached tokens from prompt caching
62    #[serde(default)]
63    pub cached_tokens: usize,
64    /// Audio input tokens, defaulting null or missing values to zero.
65    /// Zero is omitted from serialization. [`Usage::to_normalized`] uses the
66    /// reported total to determine whether audio is additional to prompt tokens.
67    #[serde(
68        default,
69        deserialize_with = "json_utils::null_or_default",
70        skip_serializing_if = "is_zero"
71    )]
72    pub audio_tokens: usize,
73    /// Tokens written to cache on this call. `None` means unreported, not zero.
74    #[serde(default, skip_serializing_if = "Option::is_none")]
75    pub cache_write_tokens: Option<usize>,
76}
77
78/// Whether a counter is absent-as-zero, for `skip_serializing_if`.
79fn is_zero(value: &usize) -> bool {
80    *value == 0
81}
82
83#[derive(Clone, Copy, Debug, Deserialize, Serialize, Default)]
84pub struct CompletionTokensDetails {
85    /// Reasoning tokens reported by reasoning-capable providers.
86    #[serde(default)]
87    pub reasoning_tokens: usize,
88}
89
90#[derive(Clone, Copy, Debug, Deserialize, Serialize)]
91pub struct Usage {
92    pub prompt_tokens: usize,
93    #[serde(default, skip_serializing_if = "Option::is_none")]
94    pub completion_tokens: Option<usize>,
95    pub total_tokens: usize,
96    // Not aliased to Mistral's singular `prompt_token_details`: Mistral's
97    // embeddings reply carries *both* keys (the singular always `null`), and
98    // an alias makes serde reject the document as a duplicate field.
99    #[serde(skip_serializing_if = "Option::is_none")]
100    pub prompt_tokens_details: Option<PromptTokensDetails>,
101    #[serde(default, skip_serializing_if = "Option::is_none")]
102    pub completion_tokens_details: Option<CompletionTokensDetails>,
103    /// Mistral's top-level cached-prompt count, reported beside (or instead
104    /// of) `prompt_tokens_details.cached_tokens`.
105    #[serde(default, skip_serializing_if = "Option::is_none")]
106    pub num_cached_tokens: Option<u64>,
107    #[serde(default, skip_serializing_if = "Option::is_none")]
108    pub queue_time: Option<f64>,
109    #[serde(default, skip_serializing_if = "Option::is_none")]
110    pub prompt_time: Option<f64>,
111    #[serde(default, skip_serializing_if = "Option::is_none")]
112    pub completion_time: Option<f64>,
113    #[serde(default, skip_serializing_if = "Option::is_none")]
114    pub total_time: Option<f64>,
115}
116
117impl Usage {
118    pub fn new() -> Self {
119        Self {
120            prompt_tokens: 0,
121            completion_tokens: None,
122            total_tokens: 0,
123            prompt_tokens_details: None,
124            completion_tokens_details: None,
125            num_cached_tokens: None,
126            queue_time: None,
127            prompt_time: None,
128            completion_time: None,
129            total_time: None,
130        }
131    }
132}
133
134impl Default for Usage {
135    fn default() -> Self {
136        Self::new()
137    }
138}
139
140impl fmt::Display for Usage {
141    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
142        let Usage {
143            prompt_tokens,
144            total_tokens,
145            ..
146        } = self;
147        write!(
148            f,
149            "Prompt tokens: {prompt_tokens} Total tokens: {total_tokens}"
150        )
151    }
152}
153
154impl From<&Usage> for crate::completion::Usage {
155    fn from(value: &Usage) -> crate::completion::Usage {
156        value.to_normalized()
157    }
158}
159
160impl From<Usage> for crate::completion::Usage {
161    fn from(value: Usage) -> crate::completion::Usage {
162        value.to_normalized()
163    }
164}
165
166impl Usage {
167    /// Return prompt tokens plus audio only when that sum and output match the total.
168    /// Missing output counts are treated as zero for this comparison.
169    fn input_tokens(&self) -> usize {
170        let audio = self
171            .prompt_tokens_details
172            .map_or(0, |details| details.audio_tokens);
173        let beside = self.prompt_tokens.saturating_add(audio);
174        let accounted = beside.saturating_add(self.completion_tokens.unwrap_or(0));
175        if audio != 0 && accounted == self.total_tokens {
176            beside
177        } else {
178            self.prompt_tokens
179        }
180    }
181
182    /// Normalize token accounting, deriving absent output counts from the total.
183    /// Cached input prefers prompt details and falls back to `num_cached_tokens`.
184    pub fn to_normalized(&self) -> crate::completion::Usage {
185        let input_tokens = self.input_tokens();
186        let details = self.prompt_tokens_details.as_ref();
187        crate::completion::Usage {
188            input_tokens: Some(input_tokens as u64),
189            // Gateways that omit `completion_tokens` still send the total, so
190            // the completion count is the remainder.
191            output_tokens: Some(
192                self.completion_tokens
193                    .unwrap_or_else(|| self.total_tokens.saturating_sub(input_tokens))
194                    as u64,
195            ),
196            total_tokens: Some(self.total_tokens as u64),
197            cached_input_tokens: details
198                .map(|d| d.cached_tokens as u64)
199                .or(self.num_cached_tokens),
200            cache_creation_input_tokens: details
201                .and_then(|d| d.cache_write_tokens)
202                .map(|tokens| tokens as u64),
203            reasoning_tokens: self
204                .completion_tokens_details
205                .as_ref()
206                .map(|d| d.reasoning_tokens as u64),
207            ..Default::default()
208        }
209    }
210}
211
212#[cfg(test)]
213mod tests;