Skip to main content

synapse/embeddings/
mod.rs

1//! Embeddings engine: OpenAI-shaped types, the provider trait, and batch splitting.
2use crate::error::GatewayError;
3use async_trait::async_trait;
4use serde::{Deserialize, Serialize};
5
6pub mod openai;
7pub mod vertex;
8
9/// One provider call's output: a vector per input (input order) + input tokens.
10#[derive(Debug, Clone)]
11pub struct EmbedOut {
12    pub vectors: Vec<Vec<f32>>,
13    pub input_tokens: u64,
14}
15
16#[async_trait]
17pub trait EmbeddingProvider: Send + Sync {
18    /// Embed `inputs`, pinning output to `dims`. One vector per input, in order.
19    async fn embed(
20        &self,
21        model: &str,
22        inputs: &[String],
23        dims: u32,
24    ) -> Result<EmbedOut, GatewayError>;
25}
26
27/// OpenAI-shaped request. `input` accepts a single string or an array.
28#[derive(Debug, Clone, Deserialize)]
29pub struct EmbeddingRequest {
30    pub input: EmbeddingInput,
31    pub model: String,
32    #[serde(default)]
33    pub dimensions: Option<u32>,
34}
35
36#[derive(Debug, Clone, Deserialize)]
37#[serde(untagged)]
38pub enum EmbeddingInput {
39    One(String),
40    Many(Vec<String>),
41}
42
43impl EmbeddingInput {
44    pub fn into_vec(self) -> Vec<String> {
45        match self {
46            EmbeddingInput::One(s) => vec![s],
47            EmbeddingInput::Many(v) => v,
48        }
49    }
50}
51
52#[derive(Debug, Clone, Serialize)]
53pub struct EmbeddingResponse {
54    pub object: &'static str, // "list"
55    pub data: Vec<EmbeddingData>,
56    pub model: String,
57    pub usage: EmbeddingUsage,
58}
59
60#[derive(Debug, Clone, Serialize)]
61pub struct EmbeddingData {
62    pub object: &'static str, // "embedding"
63    pub index: usize,
64    pub embedding: Vec<f32>,
65}
66
67#[derive(Debug, Clone, Serialize)]
68pub struct EmbeddingUsage {
69    pub prompt_tokens: u64,
70    pub total_tokens: u64,
71}
72
73/// Build an OpenAI-shaped `EmbeddingResponse` from aggregated provider output.
74/// One `EmbeddingData` per vector (dense `index`); usage is input-token only.
75pub fn build_response(model: String, out: EmbedOut) -> EmbeddingResponse {
76    let data = out
77        .vectors
78        .into_iter()
79        .enumerate()
80        .map(|(index, embedding)| EmbeddingData {
81            object: "embedding",
82            index,
83            embedding,
84        })
85        .collect();
86    EmbeddingResponse {
87        object: "list",
88        data,
89        model,
90        usage: EmbeddingUsage {
91            prompt_tokens: out.input_tokens,
92            total_tokens: out.input_tokens,
93        },
94    }
95}
96
97/// Split `inputs` into contiguous batches no larger than `limit`, preserving order.
98pub fn split_batches(inputs: &[String], limit: usize) -> Vec<&[String]> {
99    if inputs.is_empty() {
100        return vec![];
101    }
102    inputs.chunks(limit).collect()
103}
104
105#[cfg(test)]
106mod tests {
107    use super::*;
108
109    #[test]
110    fn split_batches_preserves_order_and_limit() {
111        let v: Vec<String> = (0..5).map(|i| i.to_string()).collect();
112        let batches = split_batches(&v, 2);
113        assert_eq!(batches.len(), 3);
114        assert_eq!(batches[0], &["0".to_string(), "1".to_string()]);
115        assert_eq!(batches[2], &["4".to_string()]);
116        assert!(split_batches(&[], 2).is_empty());
117    }
118
119    #[test]
120    fn input_into_vec_handles_one_and_many() {
121        assert_eq!(
122            EmbeddingInput::One("a".into()).into_vec(),
123            vec!["a".to_string()]
124        );
125        assert_eq!(
126            EmbeddingInput::Many(vec!["a".into(), "b".into()])
127                .into_vec()
128                .len(),
129            2
130        );
131    }
132}