synapse/embeddings/
mod.rs1use crate::error::GatewayError;
3use async_trait::async_trait;
4use serde::{Deserialize, Serialize};
5
6pub mod openai;
7pub mod vertex;
8
9#[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 async fn embed(
20 &self,
21 model: &str,
22 inputs: &[String],
23 dims: u32,
24 ) -> Result<EmbedOut, GatewayError>;
25}
26
27#[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, 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, 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
73pub 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
97pub 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}