rig_core/providers/openai/
embedding.rs1use crate::json_utils;
2use serde::{Deserialize, Serialize};
3use std::fmt;
4
5pub const TEXT_EMBEDDING_3_LARGE: &str = "text-embedding-3-large";
7pub const TEXT_EMBEDDING_3_SMALL: &str = "text-embedding-3-small";
9pub 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#[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#[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
50pub(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 #[serde(default)]
63 pub cached_tokens: usize,
64 #[serde(
68 default,
69 deserialize_with = "json_utils::null_or_default",
70 skip_serializing_if = "is_zero"
71 )]
72 pub audio_tokens: usize,
73 #[serde(default, skip_serializing_if = "Option::is_none")]
75 pub cache_write_tokens: Option<usize>,
76}
77
78fn is_zero(value: &usize) -> bool {
80 *value == 0
81}
82
83#[derive(Clone, Copy, Debug, Deserialize, Serialize, Default)]
84pub struct CompletionTokensDetails {
85 #[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 #[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 #[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 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 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 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;