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