Skip to main content

synapse/
vertex_native.rs

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