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