1use 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
15const 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 region: String,
26 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 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 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 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 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 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 struct St<S> {
252 inner: S,
253 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 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 }
322 }
323 }
324 });
325
326 Ok(items)
327 }
328}
329
330fn 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
351fn 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 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 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 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
462fn is_thought_part(p: &Value) -> bool {
467 p.get("thought").and_then(Value::as_bool).unwrap_or(false)
468}
469
470fn 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
499fn map_vertex_finish(s: &str) -> FinishReason {
501 match s {
502 "MAX_TOKENS" => FinishReason::Length,
503 "STOP" => FinishReason::Stop,
504 _ => FinishReason::Stop,
505 }
506}
507
508pub 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 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 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 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 assert!(
1018 matches!(err, crate::error::GatewayError::BadRequest(_)),
1019 "expected BadRequest, got {err:?}"
1020 );
1021 }
1022}