Skip to main content

laminar_db/ai/backends/
openai.rs

1//! OpenAI-compatible remote provider (OpenAI, Azure, vLLM, local servers).
2//! Chat completions for discriminative/generative tasks; the embeddings endpoint
3//! for `ai_embed` (the only provider that supports it).
4
5use 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
20/// OpenAI-compatible HTTP provider.
21pub struct OpenAiProvider {
22    client: reqwest::Client,
23    base_url: String,
24    api_key: String,
25    max_concurrency: usize,
26}
27
28impl OpenAiProvider {
29    /// Build an OpenAI-compatible provider.
30    ///
31    /// # Errors
32    ///
33    /// Returns [`ProviderError::Transport`] if the HTTP client cannot be built.
34    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        // Build owned bodies first so the async closure captures by value, not by reference.
55        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// --- wire shapes ---
131
132#[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
201/// Extract the first choice's text and token usage. Cost is left at zero;
202/// token-to-dollar conversion needs a per-model price table.
203fn 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
213/// Sort embeddings by `index` (API may return out of order) and collect usage.
214fn 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}