Skip to main content

synapse/embeddings/
openai.rs

1//! OpenAI-compatible embeddings adapter (`POST {base}/embeddings`).
2use std::time::Duration;
3
4use serde::{Deserialize, Serialize};
5
6use crate::embeddings::{EmbedOut, EmbeddingProvider};
7use crate::error::GatewayError;
8
9pub const OPENAI_EMBED_BATCH: usize = 2048;
10
11pub struct OpenAiEmbedder {
12    base_url: String,
13    api_key: String,
14    client: reqwest::Client,
15}
16
17impl OpenAiEmbedder {
18    pub fn new(base_url: String, api_key: String, timeout: Duration) -> Self {
19        let client = reqwest::Client::builder()
20            .timeout(timeout)
21            .build()
22            .expect("reqwest client");
23        Self {
24            base_url: base_url.trim_end_matches('/').to_string(),
25            api_key,
26            client,
27        }
28    }
29}
30
31#[derive(Serialize)]
32struct EmbedReq<'a> {
33    input: &'a [String],
34    model: &'a str,
35    dimensions: u32,
36}
37
38#[derive(Deserialize)]
39struct RespDatum {
40    index: usize,
41    embedding: Vec<f32>,
42}
43
44#[derive(Deserialize)]
45struct RespUsage {
46    #[serde(default)]
47    total_tokens: u64,
48}
49
50#[derive(Deserialize)]
51struct EmbedResp {
52    data: Vec<RespDatum>,
53    #[serde(default)]
54    usage: Option<RespUsage>,
55}
56
57pub fn parse_openai_response(raw: serde_json::Value) -> Result<EmbedOut, GatewayError> {
58    let mut parsed: EmbedResp =
59        serde_json::from_value(raw).map_err(|e| GatewayError::Upstream {
60            status: 502,
61            body: format!("openai embed parse: {e}"),
62        })?;
63    parsed.data.sort_by_key(|d| d.index);
64    let input_tokens = parsed.usage.map(|u| u.total_tokens).unwrap_or(0);
65    let vectors = parsed.data.into_iter().map(|d| d.embedding).collect();
66    Ok(EmbedOut {
67        vectors,
68        input_tokens,
69    })
70}
71
72#[async_trait::async_trait]
73impl EmbeddingProvider for OpenAiEmbedder {
74    async fn embed(
75        &self,
76        model: &str,
77        inputs: &[String],
78        dims: u32,
79    ) -> Result<EmbedOut, GatewayError> {
80        let resp = self
81            .client
82            .post(format!("{}/embeddings", self.base_url))
83            .bearer_auth(&self.api_key)
84            .json(&EmbedReq {
85                input: inputs,
86                model,
87                dimensions: dims,
88            })
89            .send()
90            .await
91            .map_err(|e| GatewayError::Upstream {
92                status: 502,
93                body: format!("openai embed send: {e}"),
94            })?;
95        if !resp.status().is_success() {
96            let status = resp.status().as_u16();
97            let body = resp.text().await.unwrap_or_default();
98            return Err(GatewayError::Upstream { status, body });
99        }
100        let raw = resp
101            .json::<serde_json::Value>()
102            .await
103            .map_err(|e| GatewayError::Upstream {
104                status: 502,
105                body: format!("openai embed body: {e}"),
106            })?;
107        parse_openai_response(raw)
108    }
109}
110
111#[cfg(test)]
112mod tests {
113    use super::*;
114
115    #[test]
116    fn parses_data_sorted_by_index_and_usage() {
117        let raw = serde_json::json!({
118            "data": [
119                { "index": 1, "embedding": [0.3, 0.4] },
120                { "index": 0, "embedding": [0.1, 0.2] }
121            ],
122            "usage": { "total_tokens": 9 }
123        });
124        let out = parse_openai_response(raw).unwrap();
125        assert_eq!(out.vectors[0], vec![0.1, 0.2]); // re-ordered by index
126        assert_eq!(out.vectors[1], vec![0.3, 0.4]);
127        assert_eq!(out.input_tokens, 9);
128    }
129}