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