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