Skip to main content

synapse/embeddings/
vertex.rs

1//! Native Vertex AI embeddings via the publisher `:predict` endpoint.
2use 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(&region);
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        // Mirror the chat lane (`src/vertex_native.rs`): fetch the cached bearer
134        // via `VertexAuth::token()` (returns `Result<String, String>`).
135        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}