synapse/embeddings/
openai.rs1use 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]); assert_eq!(out.vectors[1], vec![0.3, 0.4]);
127 assert_eq!(out.input_tokens, 9);
128 }
129}