synapse/embeddings/
vertex.rs1use std::sync::Arc;
3use std::time::Duration;
4
5use serde::{Deserialize, Serialize};
6
7use crate::embeddings::{EmbedOut, EmbeddingProvider};
8use crate::error::GatewayError;
9use crate::providers::vertex_auth::VertexAuth;
10
11pub const VERTEX_EMBED_BATCH: usize = 250;
12
13pub struct VertexEmbedder {
14 auth: Arc<VertexAuth>,
15 project: String,
16 region: String,
17 endpoint_base: String,
18 client: reqwest::Client,
19}
20
21impl VertexEmbedder {
22 pub fn new(auth: Arc<VertexAuth>, project: String, region: String, timeout: Duration) -> Self {
23 let endpoint_base = if region == "global" {
24 "https://aiplatform.googleapis.com".to_string()
25 } else {
26 format!("https://{region}-aiplatform.googleapis.com")
27 };
28 let client = reqwest::Client::builder()
29 .timeout(timeout)
30 .build()
31 .expect("reqwest client");
32 Self {
33 auth,
34 project,
35 region,
36 endpoint_base,
37 client,
38 }
39 }
40
41 fn predict_url(&self, model: &str) -> String {
42 format!(
43 "{}/v1/projects/{}/locations/{}/publishers/google/models/{}:predict",
44 self.endpoint_base, self.project, self.region, model
45 )
46 }
47}
48
49#[derive(Serialize)]
50struct PredictInstance {
51 content: String,
52}
53
54#[derive(Serialize)]
55struct PredictParams {
56 #[serde(rename = "outputDimensionality")]
57 output_dimensionality: u32,
58}
59
60#[derive(Serialize)]
61pub struct PredictBody {
62 instances: Vec<PredictInstance>,
63 parameters: PredictParams,
64}
65
66pub fn build_predict_body(inputs: &[String], dims: u32) -> PredictBody {
67 PredictBody {
68 instances: inputs
69 .iter()
70 .map(|c| PredictInstance { content: c.clone() })
71 .collect(),
72 parameters: PredictParams {
73 output_dimensionality: dims,
74 },
75 }
76}
77
78#[derive(Deserialize)]
79struct RespStats {
80 #[serde(default)]
81 token_count: u64,
82}
83
84#[derive(Deserialize)]
85struct RespEmbeddings {
86 values: Vec<f32>,
87 #[serde(default)]
88 statistics: Option<RespStats>,
89}
90
91#[derive(Deserialize)]
92struct RespPrediction {
93 embeddings: RespEmbeddings,
94}
95
96#[derive(Deserialize)]
97struct PredictResp {
98 predictions: Vec<RespPrediction>,
99}
100
101pub fn parse_predict_response(raw: serde_json::Value) -> Result<EmbedOut, GatewayError> {
102 let parsed: PredictResp = serde_json::from_value(raw).map_err(|e| GatewayError::Upstream {
103 status: 502,
104 body: format!("vertex embed parse: {e}"),
105 })?;
106 let input_tokens = parsed
107 .predictions
108 .iter()
109 .map(|p| {
110 p.embeddings
111 .statistics
112 .as_ref()
113 .map(|s| s.token_count)
114 .unwrap_or(0)
115 })
116 .sum();
117 let vectors = parsed
118 .predictions
119 .into_iter()
120 .map(|p| p.embeddings.values)
121 .collect();
122 Ok(EmbedOut {
123 vectors,
124 input_tokens,
125 })
126}
127
128#[async_trait::async_trait]
129impl EmbeddingProvider for VertexEmbedder {
130 async fn embed(
131 &self,
132 model: &str,
133 inputs: &[String],
134 dims: u32,
135 ) -> Result<EmbedOut, GatewayError> {
136 let token = self
139 .auth
140 .token()
141 .await
142 .map_err(|e| GatewayError::Upstream {
143 status: 502,
144 body: format!("vertex auth: {e}"),
145 })?;
146 let resp = self
147 .client
148 .post(self.predict_url(model))
149 .bearer_auth(token)
150 .json(&build_predict_body(inputs, dims))
151 .send()
152 .await
153 .map_err(|e| GatewayError::Upstream {
154 status: 502,
155 body: format!("vertex embed send: {e}"),
156 })?;
157 if !resp.status().is_success() {
158 let status = resp.status().as_u16();
159 let body = resp.text().await.unwrap_or_default();
160 return Err(GatewayError::Upstream { status, body });
161 }
162 let raw = resp
163 .json::<serde_json::Value>()
164 .await
165 .map_err(|e| GatewayError::Upstream {
166 status: 502,
167 body: format!("vertex embed body: {e}"),
168 })?;
169 parse_predict_response(raw)
170 }
171}
172
173#[cfg(test)]
174mod tests {
175 use super::*;
176 #[test]
177 fn builds_request_body_with_output_dimensionality() {
178 let body = build_predict_body(&["hello".to_string(), "world".to_string()], 768);
179 let v = serde_json::to_value(&body).unwrap();
180 assert_eq!(v["instances"][0]["content"], "hello");
181 assert_eq!(v["instances"][1]["content"], "world");
182 assert_eq!(v["parameters"]["outputDimensionality"], 768);
183 }
184 #[test]
185 fn parses_predictions_and_tokens() {
186 let raw = serde_json::json!({
187 "predictions": [
188 { "embeddings": { "values": [0.1, 0.2], "statistics": { "token_count": 3 } } },
189 { "embeddings": { "values": [0.3, 0.4], "statistics": { "token_count": 4 } } }
190 ]
191 });
192 let out = parse_predict_response(raw).unwrap();
193 assert_eq!(out.vectors.len(), 2);
194 assert_eq!(out.vectors[0], vec![0.1, 0.2]);
195 assert_eq!(out.input_tokens, 7);
196 }
197}