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 parts = crate::routing::content_parts::content_to_vertex_parts(&m.content);
303                contents.push(json!({ "role": vrole, "parts": parts }));
304            }
305        }
306    }
307    // Attach media to the last user content (or a fresh one) if present.
308    if !media_parts.is_empty() {
309        if let Some(last) = contents.iter_mut().rev().find(|c| c["role"] == "user") {
310            if let Some(arr) = last["parts"].as_array_mut() {
311                arr.extend(media_parts);
312            }
313        } else {
314            contents.push(json!({ "role": "user", "parts": media_parts }));
315        }
316    }
317
318    let mut body = json!({ "contents": contents });
319
320    if let Some(cache) = &ext.cached_content {
321        body["cachedContent"] = json!(cache);
322    }
323    if let Some(schema) = &ext.response_schema {
324        body["generationConfig"] = json!({
325            "responseMimeType": "application/json",
326            "responseSchema": schema,
327        });
328    }
329    if let Some(t) = req.temperature {
330        body["generationConfig"]["temperature"] = json!(t);
331    }
332    if let Some(tools) = &req.tools {
333        let decls = tools.iter().filter_map(|t| {
334            let f = t.get("function")?;
335            Some(json!({
336                "name": f.get("name")?.as_str()?,
337                "description": f.get("description").and_then(|d| d.as_str()).unwrap_or(""),
338                "parameters": f.get("parameters").cloned().unwrap_or_else(|| json!({"type":"object"})),
339            }))
340        }).collect::<Vec<_>>();
341        if !decls.is_empty() {
342            body["tools"] = json!([{ "functionDeclarations": decls }]);
343        }
344    }
345    if let Some(choice) = &req.tool_choice {
346        let mode = match choice {
347            Value::String(s) if s == "none" => "NONE",
348            Value::String(s) if s == "required" => "ANY",
349            Value::String(_) => "AUTO",
350            Value::Object(_) => "ANY",
351            _ => "AUTO",
352        };
353        body["toolConfig"] = json!({ "functionCallingConfig": { "mode": mode } });
354    }
355
356    body
357}
358
359/// Map a Vertex response into the shared `Completion`, extracting usage.
360fn parse_response(
361    provider: &str,
362    model: &str,
363    v: &Value,
364) -> Result<Completion, crate::error::GatewayError> {
365    let content = v["candidates"][0]["content"]["parts"]
366        .as_array()
367        .map(|parts| {
368            parts
369                .iter()
370                .filter_map(|p| p["text"].as_str())
371                .collect::<Vec<_>>()
372                .join("")
373        })
374        .unwrap_or_default();
375    let usage = &v["usageMetadata"];
376    Ok(Completion {
377        provider: provider.to_string(),
378        model: model.to_string(),
379        content,
380        tool_calls: Vec::new(),
381        finish_reason: FinishReason::Stop,
382        input_tokens: usage["promptTokenCount"].as_u64().unwrap_or(0),
383        output_tokens: usage["candidatesTokenCount"].as_u64().unwrap_or(0),
384    })
385}
386
387/// Map a Vertex finishReason string to our FinishReason.
388fn map_vertex_finish(s: &str) -> FinishReason {
389    match s {
390        "MAX_TOKENS" => FinishReason::Length,
391        "STOP" => FinishReason::Stop,
392        _ => FinishReason::Stop,
393    }
394}
395
396/// Convert one Vertex stream chunk into zero or more `StreamItem`s.
397/// `tool_index` is a running counter the caller threads across the whole stream
398/// so synthesized ids (`call_{n}`) and indices stay stable and unique.
399pub fn vertex_chunk_to_items(chunk: &Value, tool_index: &mut u32) -> Vec<StreamItem> {
400    let mut out = Vec::new();
401    if let Some(parts) = chunk["candidates"][0]["content"]["parts"].as_array() {
402        for p in parts {
403            if let Some(text) = p["text"].as_str() {
404                if !text.is_empty() {
405                    out.push(StreamItem::Delta(text.to_string()));
406                }
407            } else if let Some(fc) = p.get("functionCall") {
408                let name = fc["name"].as_str().unwrap_or_default().to_string();
409                let args = fc.get("args").cloned().unwrap_or_else(|| json!({}));
410                let i = *tool_index;
411                *tool_index += 1;
412                out.push(StreamItem::ToolCallDelta {
413                    index: i,
414                    id: Some(format!("call_{i}")),
415                    name: Some(name),
416                    args_fragment: args.to_string(),
417                });
418            }
419        }
420    }
421    // Emit the terminal Done only when Vertex signals end-of-turn via
422    // `finishReason`. Gemini includes (cumulative) `usageMetadata` on
423    // intermediate chunks too, so gating on usage presence would emit a
424    // spurious Done per chunk — harmless for the buffered accumulator but it
425    // would inject premature `finish_reason` chunks into the SSE stream.
426    // The tool-call finish override (functionCall ends with finishReason STOP)
427    // is applied by the stream driver via its `saw_tool` flag.
428    if let Some(finish) = chunk["candidates"][0]["finishReason"].as_str() {
429        let usage = &chunk["usageMetadata"];
430        out.push(StreamItem::Done {
431            input_tokens: usage["promptTokenCount"].as_u64().unwrap_or(0),
432            output_tokens: usage["candidatesTokenCount"].as_u64().unwrap_or(0),
433            finish_reason: map_vertex_finish(finish),
434        });
435    }
436    out
437}
438
439#[cfg(test)]
440mod tests {
441    use super::*;
442
443    fn req_with(ext: VertexExt) -> ChatRequest {
444        let mut r: ChatRequest = serde_json::from_value(serde_json::json!({
445            "model": "gemini-pro", "messages": [{"role": "user", "content": "describe"}]
446        }))
447        .unwrap();
448        r.vertex = Some(ext);
449        r
450    }
451
452    #[test]
453    fn payload_includes_cached_content_and_schema_and_media() {
454        let ext = VertexExt {
455            cached_content: Some("cachedContents/abc".into()),
456            media_uris: Some(vec!["gs://bucket/v.mp4".into()]),
457            response_schema: Some(serde_json::json!({"type": "object"})),
458        };
459        let body = build_payload(&req_with(ext.clone()), &ext);
460        assert_eq!(
461            body["cachedContent"],
462            serde_json::json!("cachedContents/abc")
463        );
464        assert_eq!(
465            body["generationConfig"]["responseSchema"],
466            serde_json::json!({"type": "object"})
467        );
468        let parts = body["contents"][0]["parts"].as_array().unwrap();
469        assert!(parts
470            .iter()
471            .any(|p| p["fileData"]["fileUri"] == "gs://bucket/v.mp4"));
472    }
473
474    #[test]
475    fn parses_usage_from_vertex_response() {
476        let v = serde_json::json!({
477            "candidates": [{"content": {"parts": [{"text": "a"}, {"text": "b"}], "role": "model"}}],
478            "usageMetadata": {"promptTokenCount": 10, "candidatesTokenCount": 4}
479        });
480        let c = parse_response("vertex", "gemini-3-pro", &v).unwrap();
481        assert_eq!(c.content, "ab");
482        assert_eq!(c.input_tokens, 10);
483        assert_eq!(c.output_tokens, 4);
484    }
485
486    #[test]
487    fn payload_includes_tools_and_function_messages() {
488        let r: ChatRequest = serde_json::from_value(serde_json::json!({
489            "model": "gemini-pro",
490            "messages": [
491                {"role": "user", "content": "weather?"},
492                {"role": "assistant", "content": null, "tool_calls": [
493                    {"id": "call_0", "type": "function", "function": {"name": "get_weather", "arguments": "{\"c\":\"SF\"}"}}]},
494                {"role": "tool", "tool_call_id": "call_0", "content": "21C"}
495            ],
496            "tools": [{"type": "function", "function": {"name": "get_weather",
497                "description": "Lookup", "parameters": {"type": "object"}}}],
498            "tool_choice": "auto"
499        })).unwrap();
500        let body = build_payload(&r, &VertexExt::default());
501        assert_eq!(
502            body["tools"][0]["functionDeclarations"][0]["name"],
503            "get_weather"
504        );
505        assert_eq!(body["toolConfig"]["functionCallingConfig"]["mode"], "AUTO");
506        let contents = body["contents"].as_array().unwrap();
507        assert!(contents
508            .iter()
509            .any(|c| c["parts"][0]["functionCall"]["name"] == "get_weather"));
510        assert!(contents
511            .iter()
512            .any(|c| c["parts"][0]["functionResponse"]["name"].is_string()));
513    }
514
515    #[test]
516    fn parses_vertex_chunk_text_and_functioncall() {
517        use crate::routing::stream::{FinishReason, StreamItem};
518        let mut idx = 0u32;
519
520        let text_chunk = serde_json::json!({
521            "candidates": [{"content": {"role": "model", "parts": [{"text": "Hi"}]}}]
522        });
523        let items = vertex_chunk_to_items(&text_chunk, &mut idx);
524        assert_eq!(items, vec![StreamItem::Delta("Hi".into())]);
525
526        let fc_chunk = serde_json::json!({
527            "candidates": [{"content": {"role": "model", "parts": [
528                {"functionCall": {"name": "get_weather", "args": {"c": "SF"}}}]}}]
529        });
530        let items = vertex_chunk_to_items(&fc_chunk, &mut idx);
531        assert_eq!(
532            items,
533            vec![StreamItem::ToolCallDelta {
534                index: 0,
535                id: Some("call_0".into()),
536                name: Some("get_weather".into()),
537                args_fragment: "{\"c\":\"SF\"}".into(),
538            }]
539        );
540
541        let final_chunk = serde_json::json!({
542            "candidates": [{"finishReason": "STOP", "content": {"role": "model", "parts": []}}],
543            "usageMetadata": {"promptTokenCount": 7, "candidatesTokenCount": 3}
544        });
545        let items = vertex_chunk_to_items(&final_chunk, &mut idx);
546        assert_eq!(
547            items,
548            vec![StreamItem::Done {
549                input_tokens: 7,
550                output_tokens: 3,
551                finish_reason: FinishReason::Stop
552            }]
553        );
554    }
555
556    #[tokio::test]
557    async fn generate_posts_to_vertex_url_with_bearer_and_parses() {
558        use wiremock::matchers::{header, method, path};
559        use wiremock::{Mock, MockServer, ResponseTemplate};
560
561        let mock = MockServer::start().await;
562        Mock::given(method("POST"))
563            .and(path("/v1/projects/p/locations/global/publishers/google/models/gemini-3-pro:generateContent"))
564            .and(header("authorization", "Bearer test-token"))
565            .respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({
566                "candidates": [{"content": {"parts": [{"text": "ok"}], "role": "model"}}],
567                "usageMetadata": {"promptTokenCount": 2, "candidatesTokenCount": 1}
568            })))
569            .mount(&mock)
570            .await;
571
572        let auth = Arc::new(VertexAuth::with_fetcher(|| {
573            Box::pin(async { Ok(("test-token".into(), Duration::from_secs(3600))) })
574        }));
575        let provider = VertexNativeProvider::new(
576            auth,
577            "p".into(),
578            "global".into(),
579            Duration::from_secs(5),
580            Some(mock.uri()),
581        );
582        let c = provider
583            .generate(
584                "gemini-3-pro",
585                &req_with(VertexExt {
586                    cached_content: Some("cachedContents/x".into()),
587                    ..Default::default()
588                }),
589            )
590            .await
591            .unwrap();
592        assert_eq!(c.content, "ok");
593        assert_eq!(c.input_tokens, 2);
594    }
595
596    #[tokio::test]
597    async fn stream_generate_yields_text_and_tool_done() {
598        use crate::routing::stream::{FinishReason, StreamItem};
599        use futures::StreamExt;
600        use wiremock::matchers::{method, path};
601        use wiremock::{Mock, MockServer, ResponseTemplate};
602
603        let sse = "data: {\"candidates\":[{\"content\":{\"role\":\"model\",\"parts\":[{\"text\":\"Hi\"}]}}]}\n\n\
604                   data: {\"candidates\":[{\"content\":{\"role\":\"model\",\"parts\":[{\"functionCall\":{\"name\":\"f\",\"args\":{\"a\":1}}}]}}]}\n\n\
605                   data: {\"candidates\":[{\"finishReason\":\"STOP\",\"content\":{\"role\":\"model\",\"parts\":[]}}],\"usageMetadata\":{\"promptTokenCount\":5,\"candidatesTokenCount\":4}}\n\n";
606        let mock = MockServer::start().await;
607        Mock::given(method("POST"))
608            .and(path("/v1/projects/p/locations/global/publishers/google/models/gemini-3-pro:streamGenerateContent"))
609            .respond_with(ResponseTemplate::new(200)
610                .insert_header("content-type", "text/event-stream")
611                .set_body_string(sse))
612            .mount(&mock)
613            .await;
614
615        let auth = Arc::new(VertexAuth::with_fetcher(|| {
616            Box::pin(async { Ok(("test-token".into(), Duration::from_secs(3600))) })
617        }));
618        let provider = VertexNativeProvider::new(
619            auth,
620            "p".into(),
621            "global".into(),
622            Duration::from_secs(5),
623            Some(mock.uri()),
624        );
625
626        let req: ChatRequest = serde_json::from_value(serde_json::json!({
627            "model":"gemini-3-pro","messages":[{"role":"user","content":"hi"}],"stream":true}))
628        .unwrap();
629        let mut stream = std::pin::pin!(provider
630            .stream_generate("gemini-3-pro", &req)
631            .await
632            .expect("starts"));
633        let mut items = Vec::new();
634        while let Some(it) = stream.next().await {
635            items.push(it.unwrap());
636        }
637
638        assert!(items
639            .iter()
640            .any(|i| matches!(i, StreamItem::Delta(t) if t == "Hi")));
641        assert!(items
642            .iter()
643            .any(|i| matches!(i, StreamItem::ToolCallDelta { name: Some(n), .. } if n == "f")));
644        assert!(matches!(
645            items.last().unwrap(),
646            StreamItem::Done {
647                input_tokens: 5,
648                output_tokens: 4,
649                finish_reason: FinishReason::ToolCalls
650            }
651        ));
652    }
653
654    #[tokio::test]
655    async fn stream_generate_maps_vertex_4xx_to_bad_request() {
656        use wiremock::matchers::{method, path};
657        use wiremock::{Mock, MockServer, ResponseTemplate};
658
659        let mock = MockServer::start().await;
660        Mock::given(method("POST"))
661            .and(path("/v1/projects/p/locations/global/publishers/google/models/gemini-3-pro:streamGenerateContent"))
662            .respond_with(ResponseTemplate::new(400).set_body_string("bad responseSchema"))
663            .mount(&mock)
664            .await;
665
666        let auth = Arc::new(VertexAuth::with_fetcher(|| {
667            Box::pin(async { Ok(("test-token".into(), Duration::from_secs(3600))) })
668        }));
669        let provider = VertexNativeProvider::new(
670            auth,
671            "p".into(),
672            "global".into(),
673            Duration::from_secs(5),
674            Some(mock.uri()),
675        );
676        let req: ChatRequest = serde_json::from_value(serde_json::json!({
677            "model":"gemini-3-pro","messages":[{"role":"user","content":"hi"}],"stream":true}))
678        .unwrap();
679
680        let err = provider
681            .stream_generate("gemini-3-pro", &req)
682            .await
683            .err()
684            .expect("4xx should be an error");
685        // A Vertex 4xx must surface as a client error (400), not a flat 502.
686        assert!(
687            matches!(err, crate::error::GatewayError::BadRequest(_)),
688            "expected BadRequest, got {err:?}"
689        );
690    }
691}