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