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;
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        // Mirror the chat lane (`src/vertex_native.rs`): fetch the cached bearer
137        // via `VertexAuth::token()` (returns `Result<String, String>`).
138        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}