Skip to main content

synapse/
vertex_native.rs

1//! Native Vertex REST lane: preserves cachedContents, gs:// media URIs, and
2//! strict responseSchema that the OpenAI-compatible standard lane cannot express.
3
4use std::sync::Arc;
5use std::time::Duration;
6
7use serde_json::{json, Value};
8
9use crate::providers::vertex_auth::VertexAuth;
10use crate::routing::executor::Completion;
11use crate::routing::request::{ChatRequest, VertexExt};
12use crate::routing::stream::{FinishReason, StreamItem};
13
14#[derive(Debug, Clone)]
15pub struct VertexNativeProvider {
16    http: reqwest::Client,
17    auth: Arc<VertexAuth>,
18    project: String,
19    region: String,
20    /// Base host; overridden in tests with a wiremock URI.
21    endpoint_base: String,
22}
23
24impl VertexNativeProvider {
25    pub fn new(
26        auth: Arc<VertexAuth>,
27        project: String,
28        region: String,
29        request_timeout: Duration,
30        endpoint_override: Option<String>,
31    ) -> Self {
32        let endpoint_base = endpoint_override.unwrap_or_else(|| {
33            if region == "global" {
34                "https://aiplatform.googleapis.com".into()
35            } else {
36                format!("https://{region}-aiplatform.googleapis.com")
37            }
38        });
39        Self {
40            http: reqwest::Client::builder()
41                .timeout(request_timeout)
42                .build()
43                .unwrap(),
44            auth,
45            project,
46            region,
47            endpoint_base,
48        }
49    }
50
51    fn generate_url(&self, model: &str) -> String {
52        format!(
53            "{}/v1/projects/{}/locations/{}/publishers/google/models/{}:generateContent",
54            self.endpoint_base, self.project, self.region, model
55        )
56    }
57
58    /// SUPERSEDED by `stream_generate` + the buffered aggregator. The live native
59    /// path (both stream and non-stream via `collect_committed`) goes through
60    /// `stream_generate`; this unary call is retained only for its endpoint/auth
61    /// tests. Note `parse_response` does NOT extract tool calls — do not re-wire
62    /// the request path back through here without restoring that.
63    pub async fn generate(
64        &self,
65        model: &str,
66        req: &ChatRequest,
67    ) -> Result<Completion, crate::error::GatewayError> {
68        let ext = req.vertex.clone().unwrap_or_default();
69        let payload = build_payload(req, &ext);
70        let token = self
71            .auth
72            .token()
73            .await
74            .map_err(|e| crate::error::GatewayError::Upstream {
75                status: 401,
76                body: format!("vertex auth: {e}"),
77            })?;
78
79        let resp = self
80            .http
81            .post(self.generate_url(model))
82            .bearer_auth(token)
83            .json(&payload)
84            .send()
85            .await
86            .map_err(|e| crate::error::GatewayError::Upstream {
87                status: 502,
88                body: e.to_string(),
89            })?;
90
91        let status = resp.status();
92        let value: Value = resp
93            .json()
94            .await
95            .map_err(|e| crate::error::GatewayError::Upstream {
96                status: status.as_u16(),
97                body: e.to_string(),
98            })?;
99        if !status.is_success() {
100            if status.is_client_error() {
101                return Err(crate::error::GatewayError::BadRequest(format!(
102                    "vertex {}: {}",
103                    status.as_u16(),
104                    value
105                )));
106            }
107            return Err(crate::error::GatewayError::Upstream {
108                status: status.as_u16(),
109                body: value.to_string(),
110            });
111        }
112        parse_response("vertex", model, &value)
113    }
114
115    fn stream_url(&self, model: &str) -> String {
116        format!(
117            "{}/v1/projects/{}/locations/{}/publishers/google/models/{}:streamGenerateContent?alt=sse",
118            self.endpoint_base, self.project, self.region, model
119        )
120    }
121
122    /// Open a Vertex SSE stream and normalize chunks into `StreamItem`s.
123    ///
124    /// The outer `Err` is the connection phase, returned as a `GatewayError` so
125    /// the handler preserves Vertex's status distinction (4xx -> 400 BadRequest,
126    /// 5xx/transport -> 502 Upstream) — the same mapping the non-stream
127    /// `generate` used. Per-item (mid-stream) errors remain `LegError`.
128    pub async fn stream_generate(
129        &self,
130        model: &str,
131        req: &ChatRequest,
132    ) -> Result<
133        impl futures::Stream<Item = Result<StreamItem, crate::routing::executor::LegError>>,
134        crate::error::GatewayError,
135    > {
136        use crate::error::GatewayError;
137        use crate::routing::executor::LegError;
138        use crate::routing::stream::FinishReason;
139        use futures::StreamExt;
140
141        let ext = req.vertex.clone().unwrap_or_default();
142        let payload = build_payload(req, &ext);
143        let token = self
144            .auth
145            .token()
146            .await
147            .map_err(|e| GatewayError::Upstream {
148                status: 401,
149                body: format!("vertex auth: {e}"),
150            })?;
151
152        let resp = self
153            .http
154            .post(self.stream_url(model))
155            .bearer_auth(token)
156            .json(&payload)
157            .send()
158            .await
159            .map_err(|e| GatewayError::Upstream {
160                status: 502,
161                body: e.to_string(),
162            })?;
163
164        let status = resp.status();
165        if !status.is_success() {
166            let body = resp.text().await.unwrap_or_default();
167            if status.is_client_error() {
168                return Err(GatewayError::BadRequest(format!(
169                    "vertex {}: {body}",
170                    status.as_u16()
171                )));
172            }
173            return Err(GatewayError::Upstream {
174                status: status.as_u16(),
175                body,
176            });
177        }
178
179        // State threaded across byte chunks via `unfold`: SSE line buffer,
180        // a running tool index, a saw_tool flag (to override the finish reason),
181        // and a queue of already-parsed items ready to yield.
182        struct St<S> {
183            inner: S,
184            // Raw byte buffer (NOT String): a multi-byte UTF-8 codepoint can be
185            // split across two `bytes_stream()` chunks, so we must not lossily
186            // decode per chunk. We decode only whole `\n\n`-terminated events,
187            // whose boundary is always on an ASCII byte.
188            buf: Vec<u8>,
189            tool_index: u32,
190            saw_tool: bool,
191            pending: std::collections::VecDeque<Result<StreamItem, LegError>>,
192        }
193        let state = St {
194            inner: Box::pin(resp.bytes_stream()),
195            buf: Vec::new(),
196            tool_index: 0,
197            saw_tool: false,
198            pending: std::collections::VecDeque::new(),
199        };
200
201        let items = futures::stream::unfold(state, |mut st| async move {
202            loop {
203                if let Some(item) = st.pending.pop_front() {
204                    return Some((item, st));
205                }
206                match st.inner.next().await {
207                    None => return None,
208                    Some(Err(e)) => return Some((Err(LegError::MidStream(e.to_string())), st)),
209                    Some(Ok(bytes)) => {
210                        st.buf.extend_from_slice(&bytes);
211                        // Drain complete SSE events (separated by a blank line).
212                        // The partial tail (possibly mid-codepoint) stays in `buf`
213                        // as raw bytes until its terminating `\n\n` arrives.
214                        while let Some(pos) = st.buf.windows(2).position(|w| w == b"\n\n") {
215                            let event_bytes: Vec<u8> = st.buf.drain(..pos + 2).collect();
216                            let event = String::from_utf8_lossy(&event_bytes);
217                            for line in event.lines() {
218                                let data = match line.strip_prefix("data:") {
219                                    Some(d) => d.trim(),
220                                    None => continue,
221                                };
222                                if data == "[DONE]" || data.is_empty() {
223                                    continue;
224                                }
225                                match serde_json::from_str::<serde_json::Value>(data) {
226                                    Ok(json) => {
227                                        for mut item in
228                                            vertex_chunk_to_items(&json, &mut st.tool_index)
229                                        {
230                                            if matches!(item, StreamItem::ToolCallDelta { .. }) {
231                                                st.saw_tool = true;
232                                            }
233                                            if let StreamItem::Done { finish_reason, .. } =
234                                                &mut item
235                                            {
236                                                if st.saw_tool {
237                                                    *finish_reason = FinishReason::ToolCalls;
238                                                }
239                                            }
240                                            st.pending.push_back(Ok(item));
241                                        }
242                                    }
243                                    Err(e) => st.pending.push_back(Err(LegError::MidStream(
244                                        format!("bad sse json: {e}"),
245                                    ))),
246                                }
247                            }
248                        }
249                        // loop back: either yield from `pending` or poll more bytes
250                    }
251                }
252            }
253        });
254
255        Ok(items)
256    }
257}
258
259/// Build a Vertex `generateContent`/`streamGenerateContent` body, threading
260/// native features (cache, media, schema) AND tool calling through.
261fn build_payload(req: &ChatRequest, ext: &VertexExt) -> Value {
262    let media_parts = ext
263        .media_uris
264        .iter()
265        .flatten()
266        .map(|uri| json!({ "fileData": { "fileUri": uri, "mimeType": "video/mp4" } }))
267        .collect::<Vec<_>>();
268
269    // One Vertex `content` per gateway message, mapping roles + tool turns.
270    let mut contents: Vec<Value> = Vec::new();
271    for m in &req.messages {
272        match m.role.as_str() {
273            "assistant" if m.tool_calls.is_some() => {
274                let parts = m
275                    .tool_calls
276                    .as_ref()
277                    .unwrap()
278                    .iter()
279                    .filter_map(|tc| {
280                        let f = tc.get("function")?;
281                        let name = f.get("name")?.as_str()?;
282                        let raw = f.get("arguments").and_then(|a| a.as_str()).unwrap_or("{}");
283                        let args: Value = serde_json::from_str(raw).unwrap_or_else(|_| json!({}));
284                        Some(json!({ "functionCall": { "name": name, "args": args } }))
285                    })
286                    .collect::<Vec<_>>();
287                contents.push(json!({ "role": "model", "parts": parts }));
288            }
289            "tool" => {
290                let name = m.name.clone().unwrap_or_default();
291                let response: Value = m
292                    .content
293                    .as_str()
294                    .map(|s| json!({ "content": s }))
295                    .unwrap_or_else(|| json!({ "content": m.content.to_string() }));
296                contents.push(json!({ "role": "user", "parts": [
297                    { "functionResponse": { "name": name, "response": response } }
298                ]}));
299            }
300            role => {
301                let vrole = if role == "assistant" { "model" } else { "user" };
302                let text = m
303                    .content
304                    .as_str()
305                    .map(str::to_string)
306                    .unwrap_or_else(|| m.content.to_string());
307                contents.push(json!({ "role": vrole, "parts": [{ "text": text }] }));
308            }
309        }
310    }
311    // Attach media to the last user content (or a fresh one) if present.
312    if !media_parts.is_empty() {
313        if let Some(last) = contents.iter_mut().rev().find(|c| c["role"] == "user") {
314            if let Some(arr) = last["parts"].as_array_mut() {
315                arr.extend(media_parts);
316            }
317        } else {
318            contents.push(json!({ "role": "user", "parts": media_parts }));
319        }
320    }
321
322    let mut body = json!({ "contents": contents });
323
324    if let Some(cache) = &ext.cached_content {
325        body["cachedContent"] = json!(cache);
326    }
327    if let Some(schema) = &ext.response_schema {
328        body["generationConfig"] = json!({
329            "responseMimeType": "application/json",
330            "responseSchema": schema,
331        });
332    }
333    if let Some(t) = req.temperature {
334        body["generationConfig"]["temperature"] = json!(t);
335    }
336    if let Some(tools) = &req.tools {
337        let decls = tools.iter().filter_map(|t| {
338            let f = t.get("function")?;
339            Some(json!({
340                "name": f.get("name")?.as_str()?,
341                "description": f.get("description").and_then(|d| d.as_str()).unwrap_or(""),
342                "parameters": f.get("parameters").cloned().unwrap_or_else(|| json!({"type":"object"})),
343            }))
344        }).collect::<Vec<_>>();
345        if !decls.is_empty() {
346            body["tools"] = json!([{ "functionDeclarations": decls }]);
347        }
348    }
349    if let Some(choice) = &req.tool_choice {
350        let mode = match choice {
351            Value::String(s) if s == "none" => "NONE",
352            Value::String(s) if s == "required" => "ANY",
353            Value::String(_) => "AUTO",
354            Value::Object(_) => "ANY",
355            _ => "AUTO",
356        };
357        body["toolConfig"] = json!({ "functionCallingConfig": { "mode": mode } });
358    }
359
360    body
361}
362
363/// Map a Vertex response into the shared `Completion`, extracting usage.
364fn parse_response(
365    provider: &str,
366    model: &str,
367    v: &Value,
368) -> Result<Completion, crate::error::GatewayError> {
369    let content = v["candidates"][0]["content"]["parts"]
370        .as_array()
371        .map(|parts| {
372            parts
373                .iter()
374                .filter_map(|p| p["text"].as_str())
375                .collect::<Vec<_>>()
376                .join("")
377        })
378        .unwrap_or_default();
379    let usage = &v["usageMetadata"];
380    Ok(Completion {
381        provider: provider.to_string(),
382        model: model.to_string(),
383        content,
384        tool_calls: Vec::new(),
385        finish_reason: FinishReason::Stop,
386        input_tokens: usage["promptTokenCount"].as_u64().unwrap_or(0),
387        output_tokens: usage["candidatesTokenCount"].as_u64().unwrap_or(0),
388    })
389}
390
391/// Map a Vertex finishReason string to our FinishReason.
392fn map_vertex_finish(s: &str) -> FinishReason {
393    match s {
394        "MAX_TOKENS" => FinishReason::Length,
395        "STOP" => FinishReason::Stop,
396        _ => FinishReason::Stop,
397    }
398}
399
400/// Convert one Vertex stream chunk into zero or more `StreamItem`s.
401/// `tool_index` is a running counter the caller threads across the whole stream
402/// so synthesized ids (`call_{n}`) and indices stay stable and unique.
403pub fn vertex_chunk_to_items(chunk: &Value, tool_index: &mut u32) -> Vec<StreamItem> {
404    let mut out = Vec::new();
405    if let Some(parts) = chunk["candidates"][0]["content"]["parts"].as_array() {
406        for p in parts {
407            if let Some(text) = p["text"].as_str() {
408                if !text.is_empty() {
409                    out.push(StreamItem::Delta(text.to_string()));
410                }
411            } else if let Some(fc) = p.get("functionCall") {
412                let name = fc["name"].as_str().unwrap_or_default().to_string();
413                let args = fc.get("args").cloned().unwrap_or_else(|| json!({}));
414                let i = *tool_index;
415                *tool_index += 1;
416                out.push(StreamItem::ToolCallDelta {
417                    index: i,
418                    id: Some(format!("call_{i}")),
419                    name: Some(name),
420                    args_fragment: args.to_string(),
421                });
422            }
423        }
424    }
425    // Emit the terminal Done only when Vertex signals end-of-turn via
426    // `finishReason`. Gemini includes (cumulative) `usageMetadata` on
427    // intermediate chunks too, so gating on usage presence would emit a
428    // spurious Done per chunk — harmless for the buffered accumulator but it
429    // would inject premature `finish_reason` chunks into the SSE stream.
430    // The tool-call finish override (functionCall ends with finishReason STOP)
431    // is applied by the stream driver via its `saw_tool` flag.
432    if let Some(finish) = chunk["candidates"][0]["finishReason"].as_str() {
433        let usage = &chunk["usageMetadata"];
434        out.push(StreamItem::Done {
435            input_tokens: usage["promptTokenCount"].as_u64().unwrap_or(0),
436            output_tokens: usage["candidatesTokenCount"].as_u64().unwrap_or(0),
437            finish_reason: map_vertex_finish(finish),
438        });
439    }
440    out
441}
442
443#[cfg(test)]
444mod tests {
445    use super::*;
446
447    fn req_with(ext: VertexExt) -> ChatRequest {
448        let mut r: ChatRequest = serde_json::from_value(serde_json::json!({
449            "model": "gemini-pro", "messages": [{"role": "user", "content": "describe"}]
450        }))
451        .unwrap();
452        r.vertex = Some(ext);
453        r
454    }
455
456    #[test]
457    fn payload_includes_cached_content_and_schema_and_media() {
458        let ext = VertexExt {
459            cached_content: Some("cachedContents/abc".into()),
460            media_uris: Some(vec!["gs://bucket/v.mp4".into()]),
461            response_schema: Some(serde_json::json!({"type": "object"})),
462        };
463        let body = build_payload(&req_with(ext.clone()), &ext);
464        assert_eq!(
465            body["cachedContent"],
466            serde_json::json!("cachedContents/abc")
467        );
468        assert_eq!(
469            body["generationConfig"]["responseSchema"],
470            serde_json::json!({"type": "object"})
471        );
472        let parts = body["contents"][0]["parts"].as_array().unwrap();
473        assert!(parts
474            .iter()
475            .any(|p| p["fileData"]["fileUri"] == "gs://bucket/v.mp4"));
476    }
477
478    #[test]
479    fn parses_usage_from_vertex_response() {
480        let v = serde_json::json!({
481            "candidates": [{"content": {"parts": [{"text": "a"}, {"text": "b"}], "role": "model"}}],
482            "usageMetadata": {"promptTokenCount": 10, "candidatesTokenCount": 4}
483        });
484        let c = parse_response("vertex", "gemini-3-pro", &v).unwrap();
485        assert_eq!(c.content, "ab");
486        assert_eq!(c.input_tokens, 10);
487        assert_eq!(c.output_tokens, 4);
488    }
489
490    #[test]
491    fn payload_includes_tools_and_function_messages() {
492        let r: ChatRequest = serde_json::from_value(serde_json::json!({
493            "model": "gemini-pro",
494            "messages": [
495                {"role": "user", "content": "weather?"},
496                {"role": "assistant", "content": null, "tool_calls": [
497                    {"id": "call_0", "type": "function", "function": {"name": "get_weather", "arguments": "{\"c\":\"SF\"}"}}]},
498                {"role": "tool", "tool_call_id": "call_0", "content": "21C"}
499            ],
500            "tools": [{"type": "function", "function": {"name": "get_weather",
501                "description": "Lookup", "parameters": {"type": "object"}}}],
502            "tool_choice": "auto"
503        })).unwrap();
504        let body = build_payload(&r, &VertexExt::default());
505        assert_eq!(
506            body["tools"][0]["functionDeclarations"][0]["name"],
507            "get_weather"
508        );
509        assert_eq!(body["toolConfig"]["functionCallingConfig"]["mode"], "AUTO");
510        let contents = body["contents"].as_array().unwrap();
511        assert!(contents
512            .iter()
513            .any(|c| c["parts"][0]["functionCall"]["name"] == "get_weather"));
514        assert!(contents
515            .iter()
516            .any(|c| c["parts"][0]["functionResponse"]["name"].is_string()));
517    }
518
519    #[test]
520    fn parses_vertex_chunk_text_and_functioncall() {
521        use crate::routing::stream::{FinishReason, StreamItem};
522        let mut idx = 0u32;
523
524        let text_chunk = serde_json::json!({
525            "candidates": [{"content": {"role": "model", "parts": [{"text": "Hi"}]}}]
526        });
527        let items = vertex_chunk_to_items(&text_chunk, &mut idx);
528        assert_eq!(items, vec![StreamItem::Delta("Hi".into())]);
529
530        let fc_chunk = serde_json::json!({
531            "candidates": [{"content": {"role": "model", "parts": [
532                {"functionCall": {"name": "get_weather", "args": {"c": "SF"}}}]}}]
533        });
534        let items = vertex_chunk_to_items(&fc_chunk, &mut idx);
535        assert_eq!(
536            items,
537            vec![StreamItem::ToolCallDelta {
538                index: 0,
539                id: Some("call_0".into()),
540                name: Some("get_weather".into()),
541                args_fragment: "{\"c\":\"SF\"}".into(),
542            }]
543        );
544
545        let final_chunk = serde_json::json!({
546            "candidates": [{"finishReason": "STOP", "content": {"role": "model", "parts": []}}],
547            "usageMetadata": {"promptTokenCount": 7, "candidatesTokenCount": 3}
548        });
549        let items = vertex_chunk_to_items(&final_chunk, &mut idx);
550        assert_eq!(
551            items,
552            vec![StreamItem::Done {
553                input_tokens: 7,
554                output_tokens: 3,
555                finish_reason: FinishReason::Stop
556            }]
557        );
558    }
559
560    #[tokio::test]
561    async fn generate_posts_to_vertex_url_with_bearer_and_parses() {
562        use wiremock::matchers::{header, method, path};
563        use wiremock::{Mock, MockServer, ResponseTemplate};
564
565        let mock = MockServer::start().await;
566        Mock::given(method("POST"))
567            .and(path("/v1/projects/p/locations/global/publishers/google/models/gemini-3-pro:generateContent"))
568            .and(header("authorization", "Bearer test-token"))
569            .respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({
570                "candidates": [{"content": {"parts": [{"text": "ok"}], "role": "model"}}],
571                "usageMetadata": {"promptTokenCount": 2, "candidatesTokenCount": 1}
572            })))
573            .mount(&mock)
574            .await;
575
576        let auth = Arc::new(VertexAuth::with_fetcher(|| {
577            Box::pin(async { Ok(("test-token".into(), Duration::from_secs(3600))) })
578        }));
579        let provider = VertexNativeProvider::new(
580            auth,
581            "p".into(),
582            "global".into(),
583            Duration::from_secs(5),
584            Some(mock.uri()),
585        );
586        let c = provider
587            .generate(
588                "gemini-3-pro",
589                &req_with(VertexExt {
590                    cached_content: Some("cachedContents/x".into()),
591                    ..Default::default()
592                }),
593            )
594            .await
595            .unwrap();
596        assert_eq!(c.content, "ok");
597        assert_eq!(c.input_tokens, 2);
598    }
599
600    #[tokio::test]
601    async fn stream_generate_yields_text_and_tool_done() {
602        use crate::routing::stream::{FinishReason, StreamItem};
603        use futures::StreamExt;
604        use wiremock::matchers::{method, path};
605        use wiremock::{Mock, MockServer, ResponseTemplate};
606
607        let sse = "data: {\"candidates\":[{\"content\":{\"role\":\"model\",\"parts\":[{\"text\":\"Hi\"}]}}]}\n\n\
608                   data: {\"candidates\":[{\"content\":{\"role\":\"model\",\"parts\":[{\"functionCall\":{\"name\":\"f\",\"args\":{\"a\":1}}}]}}]}\n\n\
609                   data: {\"candidates\":[{\"finishReason\":\"STOP\",\"content\":{\"role\":\"model\",\"parts\":[]}}],\"usageMetadata\":{\"promptTokenCount\":5,\"candidatesTokenCount\":4}}\n\n";
610        let mock = MockServer::start().await;
611        Mock::given(method("POST"))
612            .and(path("/v1/projects/p/locations/global/publishers/google/models/gemini-3-pro:streamGenerateContent"))
613            .respond_with(ResponseTemplate::new(200)
614                .insert_header("content-type", "text/event-stream")
615                .set_body_string(sse))
616            .mount(&mock)
617            .await;
618
619        let auth = Arc::new(VertexAuth::with_fetcher(|| {
620            Box::pin(async { Ok(("test-token".into(), Duration::from_secs(3600))) })
621        }));
622        let provider = VertexNativeProvider::new(
623            auth,
624            "p".into(),
625            "global".into(),
626            Duration::from_secs(5),
627            Some(mock.uri()),
628        );
629
630        let req: ChatRequest = serde_json::from_value(serde_json::json!({
631            "model":"gemini-3-pro","messages":[{"role":"user","content":"hi"}],"stream":true}))
632        .unwrap();
633        let mut stream = std::pin::pin!(provider
634            .stream_generate("gemini-3-pro", &req)
635            .await
636            .expect("starts"));
637        let mut items = Vec::new();
638        while let Some(it) = stream.next().await {
639            items.push(it.unwrap());
640        }
641
642        assert!(items
643            .iter()
644            .any(|i| matches!(i, StreamItem::Delta(t) if t == "Hi")));
645        assert!(items
646            .iter()
647            .any(|i| matches!(i, StreamItem::ToolCallDelta { name: Some(n), .. } if n == "f")));
648        assert!(matches!(
649            items.last().unwrap(),
650            StreamItem::Done {
651                input_tokens: 5,
652                output_tokens: 4,
653                finish_reason: FinishReason::ToolCalls
654            }
655        ));
656    }
657
658    #[tokio::test]
659    async fn stream_generate_maps_vertex_4xx_to_bad_request() {
660        use wiremock::matchers::{method, path};
661        use wiremock::{Mock, MockServer, ResponseTemplate};
662
663        let mock = MockServer::start().await;
664        Mock::given(method("POST"))
665            .and(path("/v1/projects/p/locations/global/publishers/google/models/gemini-3-pro:streamGenerateContent"))
666            .respond_with(ResponseTemplate::new(400).set_body_string("bad responseSchema"))
667            .mount(&mock)
668            .await;
669
670        let auth = Arc::new(VertexAuth::with_fetcher(|| {
671            Box::pin(async { Ok(("test-token".into(), Duration::from_secs(3600))) })
672        }));
673        let provider = VertexNativeProvider::new(
674            auth,
675            "p".into(),
676            "global".into(),
677            Duration::from_secs(5),
678            Some(mock.uri()),
679        );
680        let req: ChatRequest = serde_json::from_value(serde_json::json!({
681            "model":"gemini-3-pro","messages":[{"role":"user","content":"hi"}],"stream":true}))
682        .unwrap();
683
684        let err = provider
685            .stream_generate("gemini-3-pro", &req)
686            .await
687            .err()
688            .expect("4xx should be an error");
689        // A Vertex 4xx must surface as a client error (400), not a flat 502.
690        assert!(
691            matches!(err, crate::error::GatewayError::BadRequest(_)),
692            "expected BadRequest, got {err:?}"
693        );
694    }
695}