1use std::time::Duration;
6
7use async_trait::async_trait;
8use futures::stream::{StreamExt, TryStreamExt};
9use serde::{Deserialize, Serialize};
10
11use crate::ai::backends::remote::{add_usage, chat_prompt, post_json};
12use crate::ai::provider::{
13 InferenceOutputs, InferenceProvider, InferenceRequest, InferenceResponse, ProviderError, Usage,
14};
15use crate::ai::registry::Task;
16
17const REQUEST_TIMEOUT_MS: u64 = 60_000;
18const MAX_RETRIES: u32 = 2;
19
20pub struct OpenAiProvider {
22 client: reqwest::Client,
23 base_url: String,
24 api_key: String,
25 max_concurrency: usize,
26}
27
28impl OpenAiProvider {
29 pub fn new(
35 base_url: impl Into<String>,
36 api_key: impl Into<String>,
37 max_concurrency: usize,
38 ) -> Result<Self, ProviderError> {
39 let client = reqwest::Client::builder()
40 .timeout(Duration::from_millis(REQUEST_TIMEOUT_MS))
41 .build()
42 .map_err(|e| ProviderError::Transport(e.to_string()))?;
43 Ok(Self {
44 client,
45 base_url: base_url.into().trim_end_matches('/').to_string(),
46 api_key: api_key.into(),
47 max_concurrency: max_concurrency.max(1),
48 })
49 }
50
51 async fn chat(&self, request: &InferenceRequest) -> Result<InferenceResponse, ProviderError> {
52 let url = format!("{}/chat/completions", self.base_url);
53
54 let bodies: Vec<ChatBody> = request
56 .inputs
57 .iter()
58 .map(|input| {
59 let (system, user) =
60 chat_prompt(request.task, input, request.params.labels.as_deref());
61 ChatBody {
62 model: request.model.clone(),
63 messages: vec![ChatMessage::system(system), ChatMessage::user(user)],
64 }
65 })
66 .collect();
67
68 let url = &url;
69 let calls = bodies.into_iter().map(|body| async move {
70 let builder = self.client.post(url).bearer_auth(&self.api_key).json(&body);
71 let response: ChatResponse =
72 post_json(builder, MAX_RETRIES, REQUEST_TIMEOUT_MS).await?;
73 parse_chat(response)
74 });
75
76 let results: Vec<(String, Usage)> = futures::stream::iter(calls)
77 .buffered(self.max_concurrency)
78 .try_collect()
79 .await?;
80
81 let mut texts = Vec::with_capacity(results.len());
82 let mut usage = Usage::ZERO;
83 for (text, call_usage) in results {
84 texts.push(text);
85 usage = add_usage(usage, call_usage);
86 }
87 Ok(InferenceResponse {
88 outputs: InferenceOutputs::Text(texts),
89 usage,
90 })
91 }
92
93 async fn embed(&self, request: &InferenceRequest) -> Result<InferenceResponse, ProviderError> {
94 let url = format!("{}/embeddings", self.base_url);
95 let body = EmbedBody {
96 model: request.model.clone(),
97 input: request.inputs.clone(),
98 };
99 let builder = self
100 .client
101 .post(&url)
102 .bearer_auth(&self.api_key)
103 .json(&body);
104 let response: EmbedResponse = post_json(builder, MAX_RETRIES, REQUEST_TIMEOUT_MS).await?;
105 let (vectors, usage) = parse_embed(response);
106 Ok(InferenceResponse {
107 outputs: InferenceOutputs::Vectors(vectors),
108 usage,
109 })
110 }
111}
112
113#[async_trait]
114impl InferenceProvider for OpenAiProvider {
115 async fn infer_batch(
116 &self,
117 request: InferenceRequest,
118 ) -> Result<InferenceResponse, ProviderError> {
119 match request.task {
120 Task::Embed => self.embed(&request).await,
121 _ => self.chat(&request).await,
122 }
123 }
124
125 fn name(&self) -> &'static str {
126 "openai"
127 }
128}
129
130#[derive(Serialize)]
133struct ChatBody {
134 model: String,
135 messages: Vec<ChatMessage>,
136}
137
138#[derive(Serialize)]
139struct ChatMessage {
140 role: &'static str,
141 content: String,
142}
143
144impl ChatMessage {
145 fn system(content: String) -> Self {
146 Self {
147 role: "system",
148 content,
149 }
150 }
151 fn user(content: String) -> Self {
152 Self {
153 role: "user",
154 content,
155 }
156 }
157}
158
159#[derive(Deserialize)]
160struct ChatResponse {
161 choices: Vec<ChatChoice>,
162 usage: Option<TokenUsage>,
163}
164
165#[derive(Deserialize)]
166struct ChatChoice {
167 message: ChatChoiceMessage,
168}
169
170#[derive(Deserialize)]
171struct ChatChoiceMessage {
172 content: String,
173}
174
175#[derive(Serialize)]
176struct EmbedBody {
177 model: String,
178 input: Vec<String>,
179}
180
181#[derive(Deserialize)]
182struct EmbedResponse {
183 data: Vec<EmbedData>,
184 usage: Option<TokenUsage>,
185}
186
187#[derive(Deserialize)]
188struct EmbedData {
189 embedding: Vec<f32>,
190 index: usize,
191}
192
193#[derive(Deserialize)]
194struct TokenUsage {
195 #[serde(default)]
196 prompt_tokens: u64,
197 #[serde(default)]
198 completion_tokens: u64,
199}
200
201fn parse_chat(response: ChatResponse) -> Result<(String, Usage), ProviderError> {
204 let content = response
205 .choices
206 .into_iter()
207 .next()
208 .map(|c| c.message.content)
209 .ok_or_else(|| ProviderError::BadResponse("chat response had no choices".to_string()))?;
210 Ok((content, token_usage(response.usage)))
211}
212
213fn parse_embed(mut response: EmbedResponse) -> (Vec<Vec<f32>>, Usage) {
215 response.data.sort_by_key(|d| d.index);
216 let usage = token_usage(response.usage);
217 let vectors = response.data.into_iter().map(|d| d.embedding).collect();
218 (vectors, usage)
219}
220
221fn token_usage(usage: Option<TokenUsage>) -> Usage {
222 usage.map_or(Usage::ZERO, |u| Usage {
223 input_tokens: u.prompt_tokens,
224 output_tokens: u.completion_tokens,
225 cost_micros: 0,
226 })
227}
228
229#[cfg(test)]
230mod tests {
231 use super::*;
232
233 #[test]
234 fn chat_body_serializes_to_messages() {
235 let body = ChatBody {
236 model: "gpt-x".to_string(),
237 messages: vec![
238 ChatMessage::system("be terse".to_string()),
239 ChatMessage::user("hello".to_string()),
240 ],
241 };
242 let value = serde_json::to_value(&body).unwrap();
243 assert_eq!(value["model"], "gpt-x");
244 assert_eq!(value["messages"][0]["role"], "system");
245 assert_eq!(value["messages"][1]["content"], "hello");
246 }
247
248 #[test]
249 fn parse_chat_extracts_content_and_tokens() {
250 let json = r#"{
251 "choices": [{"message": {"role": "assistant", "content": "positive"}}],
252 "usage": {"prompt_tokens": 12, "completion_tokens": 1}
253 }"#;
254 let response: ChatResponse = serde_json::from_str(json).unwrap();
255 let (text, usage) = parse_chat(response).unwrap();
256 assert_eq!(text, "positive");
257 assert_eq!(usage.input_tokens, 12);
258 assert_eq!(usage.output_tokens, 1);
259 }
260
261 #[test]
262 fn parse_chat_errors_without_choices() {
263 let response: ChatResponse = serde_json::from_str(r#"{"choices": []}"#).unwrap();
264 assert!(parse_chat(response).is_err());
265 }
266
267 #[test]
268 fn parse_embed_orders_by_index() {
269 let json = r#"{
270 "data": [
271 {"embedding": [0.3, 0.4], "index": 1},
272 {"embedding": [0.1, 0.2], "index": 0}
273 ],
274 "usage": {"prompt_tokens": 5}
275 }"#;
276 let response: EmbedResponse = serde_json::from_str(json).unwrap();
277 let (vectors, usage) = parse_embed(response);
278 assert_eq!(vectors, vec![vec![0.1, 0.2], vec![0.3, 0.4]]);
279 assert_eq!(usage.input_tokens, 5);
280 }
281}