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};
13
14#[derive(Debug, Clone)]
15pub struct VertexNativeProvider {
16 http: reqwest::Client,
17 auth: Arc<VertexAuth>,
18 project: String,
19 region: String,
21 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 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 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 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 struct St<S> {
201 inner: S,
202 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 while let Some(pos) = st.buf.windows(2).position(|w| w == b"\n\n") {
233 let event_bytes: Vec<u8> = st.buf.drain(..pos + 2).collect();
234 let event = String::from_utf8_lossy(&event_bytes);
235 for line in event.lines() {
236 let data = match line.strip_prefix("data:") {
237 Some(d) => d.trim(),
238 None => continue,
239 };
240 if data == "[DONE]" || data.is_empty() {
241 continue;
242 }
243 match serde_json::from_str::<serde_json::Value>(data) {
244 Ok(json) => {
245 for mut item in
246 vertex_chunk_to_items(&json, &mut st.tool_index)
247 {
248 if matches!(item, StreamItem::ToolCallDelta { .. }) {
249 st.saw_tool = true;
250 }
251 if let StreamItem::Done { finish_reason, .. } =
252 &mut item
253 {
254 if st.saw_tool {
255 *finish_reason = FinishReason::ToolCalls;
256 }
257 }
258 st.pending.push_back(Ok(item));
259 }
260 }
261 Err(e) => st.pending.push_back(Err(LegError::MidStream(
262 format!("bad sse json: {e}"),
263 ))),
264 }
265 }
266 }
267 }
269 }
270 }
271 });
272
273 Ok(items)
274 }
275}
276
277fn build_payload(req: &ChatRequest, ext: &VertexExt) -> Value {
280 let media_parts = ext
281 .media_uris
282 .iter()
283 .flatten()
284 .map(|uri| json!({ "fileData": { "fileUri": uri, "mimeType": "video/mp4" } }))
285 .collect::<Vec<_>>();
286
287 let mut contents: Vec<Value> = Vec::new();
289 for m in &req.messages {
290 match m.role.as_str() {
291 "assistant" if m.tool_calls.is_some() => {
292 let parts = m
293 .tool_calls
294 .as_ref()
295 .unwrap()
296 .iter()
297 .filter_map(|tc| {
298 let f = tc.get("function")?;
299 let name = f.get("name")?.as_str()?;
300 let raw = f.get("arguments").and_then(|a| a.as_str()).unwrap_or("{}");
301 let args: Value = serde_json::from_str(raw).unwrap_or_else(|_| json!({}));
302 Some(json!({ "functionCall": { "name": name, "args": args } }))
303 })
304 .collect::<Vec<_>>();
305 contents.push(json!({ "role": "model", "parts": parts }));
306 }
307 "tool" => {
308 let name = m.name.clone().unwrap_or_default();
309 let response: Value = m
310 .content
311 .as_str()
312 .map(|s| json!({ "content": s }))
313 .unwrap_or_else(|| json!({ "content": m.content.to_string() }));
314 contents.push(json!({ "role": "user", "parts": [
315 { "functionResponse": { "name": name, "response": response } }
316 ]}));
317 }
318 role => {
319 let vrole = if role == "assistant" { "model" } else { "user" };
320 let parts = crate::routing::content_parts::content_to_vertex_parts(&m.content);
321 contents.push(json!({ "role": vrole, "parts": parts }));
322 }
323 }
324 }
325 if !media_parts.is_empty() {
327 if let Some(last) = contents.iter_mut().rev().find(|c| c["role"] == "user") {
328 if let Some(arr) = last["parts"].as_array_mut() {
329 arr.extend(media_parts);
330 }
331 } else {
332 contents.push(json!({ "role": "user", "parts": media_parts }));
333 }
334 }
335
336 let mut body = json!({ "contents": contents });
337
338 if let Some(cache) = &ext.cached_content {
339 body["cachedContent"] = json!(cache);
340 }
341 let mut gen_cfg = serde_json::Map::new();
345 if let Some(schema) = &ext.response_schema {
346 gen_cfg.insert("responseMimeType".into(), json!("application/json"));
347 gen_cfg.insert("responseSchema".into(), schema.clone());
348 }
349 if let Some(t) = req.temperature {
350 gen_cfg.insert("temperature".into(), json!(t));
351 }
352 if let Some(max) = req.max_tokens {
353 gen_cfg.insert("maxOutputTokens".into(), json!(max));
354 }
355 if let Some(thinking) = &ext.thinking_config {
356 gen_cfg.insert("thinkingConfig".into(), thinking.clone());
357 }
358 if !gen_cfg.is_empty() {
359 body["generationConfig"] = Value::Object(gen_cfg);
360 }
361 if let Some(tools) = &req.tools {
362 let decls = tools.iter().filter_map(|t| {
363 let f = t.get("function")?;
364 Some(json!({
365 "name": f.get("name")?.as_str()?,
366 "description": f.get("description").and_then(|d| d.as_str()).unwrap_or(""),
367 "parameters": f.get("parameters").cloned().unwrap_or_else(|| json!({"type":"object"})),
368 }))
369 }).collect::<Vec<_>>();
370 if !decls.is_empty() {
371 body["tools"] = json!([{ "functionDeclarations": decls }]);
372 }
373 }
374 if let Some(choice) = &req.tool_choice {
375 let mode = match choice {
376 Value::String(s) if s == "none" => "NONE",
377 Value::String(s) if s == "required" => "ANY",
378 Value::String(_) => "AUTO",
379 Value::Object(_) => "ANY",
380 _ => "AUTO",
381 };
382 body["toolConfig"] = json!({ "functionCallingConfig": { "mode": mode } });
383 }
384
385 body
386}
387
388fn is_thought_part(p: &Value) -> bool {
393 p.get("thought").and_then(Value::as_bool).unwrap_or(false)
394}
395
396fn parse_response(
398 provider: &str,
399 model: &str,
400 v: &Value,
401) -> Result<Completion, crate::error::GatewayError> {
402 let content = v["candidates"][0]["content"]["parts"]
403 .as_array()
404 .map(|parts| {
405 parts
406 .iter()
407 .filter(|p| !is_thought_part(p))
408 .filter_map(|p| p["text"].as_str())
409 .collect::<Vec<_>>()
410 .join("")
411 })
412 .unwrap_or_default();
413 let usage = &v["usageMetadata"];
414 Ok(Completion {
415 provider: provider.to_string(),
416 model: model.to_string(),
417 content,
418 tool_calls: Vec::new(),
419 finish_reason: FinishReason::Stop,
420 input_tokens: usage["promptTokenCount"].as_u64().unwrap_or(0),
421 output_tokens: usage["candidatesTokenCount"].as_u64().unwrap_or(0),
422 })
423}
424
425fn map_vertex_finish(s: &str) -> FinishReason {
427 match s {
428 "MAX_TOKENS" => FinishReason::Length,
429 "STOP" => FinishReason::Stop,
430 _ => FinishReason::Stop,
431 }
432}
433
434pub fn vertex_chunk_to_items(chunk: &Value, tool_index: &mut u32) -> Vec<StreamItem> {
438 let mut out = Vec::new();
439 if let Some(parts) = chunk["candidates"][0]["content"]["parts"].as_array() {
440 for p in parts {
441 if is_thought_part(p) {
442 continue;
443 }
444 if let Some(text) = p["text"].as_str() {
445 if !text.is_empty() {
446 out.push(StreamItem::Delta(text.to_string()));
447 }
448 } else if let Some(fc) = p.get("functionCall") {
449 let name = fc["name"].as_str().unwrap_or_default().to_string();
450 let args = fc.get("args").cloned().unwrap_or_else(|| json!({}));
451 let i = *tool_index;
452 *tool_index += 1;
453 out.push(StreamItem::ToolCallDelta {
454 index: i,
455 id: Some(format!("call_{i}")),
456 name: Some(name),
457 args_fragment: args.to_string(),
458 });
459 }
460 }
461 }
462 if let Some(finish) = chunk["candidates"][0]["finishReason"].as_str() {
470 let usage = &chunk["usageMetadata"];
471 out.push(StreamItem::Done {
472 input_tokens: usage["promptTokenCount"].as_u64().unwrap_or(0),
473 output_tokens: usage["candidatesTokenCount"].as_u64().unwrap_or(0),
474 finish_reason: map_vertex_finish(finish),
475 });
476 }
477 out
478}
479
480#[cfg(test)]
481mod tests {
482 use super::*;
483
484 fn req_with(ext: VertexExt) -> ChatRequest {
485 let mut r: ChatRequest = serde_json::from_value(serde_json::json!({
486 "model": "gemini-pro", "messages": [{"role": "user", "content": "describe"}]
487 }))
488 .unwrap();
489 r.vertex = Some(ext);
490 r
491 }
492
493 #[test]
494 fn payload_includes_cached_content_and_schema_and_media() {
495 let ext = VertexExt {
496 cached_content: Some("cachedContents/abc".into()),
497 media_uris: Some(vec!["gs://bucket/v.mp4".into()]),
498 response_schema: Some(serde_json::json!({"type": "object"})),
499 ..Default::default()
500 };
501 let body = build_payload(&req_with(ext.clone()), &ext);
502 assert_eq!(
503 body["cachedContent"],
504 serde_json::json!("cachedContents/abc")
505 );
506 assert_eq!(
507 body["generationConfig"]["responseSchema"],
508 serde_json::json!({"type": "object"})
509 );
510 let parts = body["contents"][0]["parts"].as_array().unwrap();
511 assert!(parts
512 .iter()
513 .any(|p| p["fileData"]["fileUri"] == "gs://bucket/v.mp4"));
514 }
515
516 #[test]
517 fn payload_includes_max_output_tokens_and_thinking_config() {
518 let ext = VertexExt {
519 thinking_config: Some(serde_json::json!({ "thinkingLevel": "low" })),
520 ..Default::default()
521 };
522 let mut req = req_with(ext.clone());
523 req.max_tokens = Some(8192);
524 let body = build_payload(&req, &ext);
525 assert_eq!(
526 body["generationConfig"]["maxOutputTokens"],
527 serde_json::json!(8192)
528 );
529 assert_eq!(
530 body["generationConfig"]["thinkingConfig"],
531 serde_json::json!({ "thinkingLevel": "low" })
532 );
533 }
534
535 #[test]
536 fn vertex_chunk_skips_thought_parts() {
537 use crate::routing::stream::StreamItem;
538 let mut idx = 0u32;
539 let chunk = serde_json::json!({
540 "candidates": [{"content": {"role": "model", "parts": [
541 {"text": "internal reasoning", "thought": true},
542 {"text": "answer"}
543 ]}}]
544 });
545 let items = vertex_chunk_to_items(&chunk, &mut idx);
546 assert_eq!(items, vec![StreamItem::Delta("answer".into())]);
547 }
548
549 #[test]
550 fn parse_response_skips_thought_parts() {
551 let v = serde_json::json!({
552 "candidates": [{"content": {"parts": [
553 {"text": "reasoning", "thought": true},
554 {"text": "real"}
555 ], "role": "model"}}],
556 "usageMetadata": {"promptTokenCount": 1, "candidatesTokenCount": 1}
557 });
558 let c = parse_response("vertex", "gemini-3-pro", &v).unwrap();
559 assert_eq!(c.content, "real");
560 }
561
562 #[test]
563 fn parses_usage_from_vertex_response() {
564 let v = serde_json::json!({
565 "candidates": [{"content": {"parts": [{"text": "a"}, {"text": "b"}], "role": "model"}}],
566 "usageMetadata": {"promptTokenCount": 10, "candidatesTokenCount": 4}
567 });
568 let c = parse_response("vertex", "gemini-3-pro", &v).unwrap();
569 assert_eq!(c.content, "ab");
570 assert_eq!(c.input_tokens, 10);
571 assert_eq!(c.output_tokens, 4);
572 }
573
574 #[test]
575 fn payload_includes_tools_and_function_messages() {
576 let r: ChatRequest = serde_json::from_value(serde_json::json!({
577 "model": "gemini-pro",
578 "messages": [
579 {"role": "user", "content": "weather?"},
580 {"role": "assistant", "content": null, "tool_calls": [
581 {"id": "call_0", "type": "function", "function": {"name": "get_weather", "arguments": "{\"c\":\"SF\"}"}}]},
582 {"role": "tool", "tool_call_id": "call_0", "content": "21C"}
583 ],
584 "tools": [{"type": "function", "function": {"name": "get_weather",
585 "description": "Lookup", "parameters": {"type": "object"}}}],
586 "tool_choice": "auto"
587 })).unwrap();
588 let body = build_payload(&r, &VertexExt::default());
589 assert_eq!(
590 body["tools"][0]["functionDeclarations"][0]["name"],
591 "get_weather"
592 );
593 assert_eq!(body["toolConfig"]["functionCallingConfig"]["mode"], "AUTO");
594 let contents = body["contents"].as_array().unwrap();
595 assert!(contents
596 .iter()
597 .any(|c| c["parts"][0]["functionCall"]["name"] == "get_weather"));
598 assert!(contents
599 .iter()
600 .any(|c| c["parts"][0]["functionResponse"]["name"].is_string()));
601 }
602
603 #[test]
604 fn parses_vertex_chunk_text_and_functioncall() {
605 use crate::routing::stream::{FinishReason, StreamItem};
606 let mut idx = 0u32;
607
608 let text_chunk = serde_json::json!({
609 "candidates": [{"content": {"role": "model", "parts": [{"text": "Hi"}]}}]
610 });
611 let items = vertex_chunk_to_items(&text_chunk, &mut idx);
612 assert_eq!(items, vec![StreamItem::Delta("Hi".into())]);
613
614 let fc_chunk = serde_json::json!({
615 "candidates": [{"content": {"role": "model", "parts": [
616 {"functionCall": {"name": "get_weather", "args": {"c": "SF"}}}]}}]
617 });
618 let items = vertex_chunk_to_items(&fc_chunk, &mut idx);
619 assert_eq!(
620 items,
621 vec![StreamItem::ToolCallDelta {
622 index: 0,
623 id: Some("call_0".into()),
624 name: Some("get_weather".into()),
625 args_fragment: "{\"c\":\"SF\"}".into(),
626 }]
627 );
628
629 let final_chunk = serde_json::json!({
630 "candidates": [{"finishReason": "STOP", "content": {"role": "model", "parts": []}}],
631 "usageMetadata": {"promptTokenCount": 7, "candidatesTokenCount": 3}
632 });
633 let items = vertex_chunk_to_items(&final_chunk, &mut idx);
634 assert_eq!(
635 items,
636 vec![StreamItem::Done {
637 input_tokens: 7,
638 output_tokens: 3,
639 finish_reason: FinishReason::Stop
640 }]
641 );
642 }
643
644 #[test]
645 fn endpoint_for_region_picks_regional_or_global_host() {
646 let auth = Arc::new(VertexAuth::with_fetcher(|| {
647 Box::pin(async { Ok(("t".into(), Duration::from_secs(3600))) })
648 }));
649 let provider = VertexNativeProvider::new(
650 auth,
651 "p".into(),
652 "global".into(),
653 Duration::from_secs(5),
654 None,
655 );
656 assert_eq!(
657 provider.endpoint_for("global"),
658 "https://aiplatform.googleapis.com"
659 );
660 assert_eq!(
661 provider.endpoint_for("us-central1"),
662 "https://us-central1-aiplatform.googleapis.com"
663 );
664 }
665
666 #[tokio::test]
667 async fn generate_uses_per_leg_region_override_in_url() {
668 use wiremock::matchers::{method, path};
669 use wiremock::{Mock, MockServer, ResponseTemplate};
670
671 let mock = MockServer::start().await;
672 Mock::given(method("POST"))
673 .and(path("/v1/projects/p/locations/us-central1/publishers/google/models/gemini-x:generateContent"))
674 .respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({
675 "candidates": [{"content": {"parts": [{"text": "ok"}], "role": "model"}}],
676 "usageMetadata": {"promptTokenCount": 1, "candidatesTokenCount": 1}
677 })))
678 .mount(&mock)
679 .await;
680
681 let auth = Arc::new(VertexAuth::with_fetcher(|| {
682 Box::pin(async { Ok(("test-token".into(), Duration::from_secs(3600))) })
683 }));
684 let provider = VertexNativeProvider::new(
686 auth,
687 "p".into(),
688 "global".into(),
689 Duration::from_secs(5),
690 Some(mock.uri()),
691 );
692 let c = provider
693 .generate(
694 "gemini-x",
695 &req_with(VertexExt::default()),
696 Some("us-central1"),
697 )
698 .await
699 .unwrap();
700 assert_eq!(c.content, "ok");
701 }
702
703 #[tokio::test]
704 async fn generate_posts_to_vertex_url_with_bearer_and_parses() {
705 use wiremock::matchers::{header, method, path};
706 use wiremock::{Mock, MockServer, ResponseTemplate};
707
708 let mock = MockServer::start().await;
709 Mock::given(method("POST"))
710 .and(path("/v1/projects/p/locations/global/publishers/google/models/gemini-3-pro:generateContent"))
711 .and(header("authorization", "Bearer test-token"))
712 .respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({
713 "candidates": [{"content": {"parts": [{"text": "ok"}], "role": "model"}}],
714 "usageMetadata": {"promptTokenCount": 2, "candidatesTokenCount": 1}
715 })))
716 .mount(&mock)
717 .await;
718
719 let auth = Arc::new(VertexAuth::with_fetcher(|| {
720 Box::pin(async { Ok(("test-token".into(), Duration::from_secs(3600))) })
721 }));
722 let provider = VertexNativeProvider::new(
723 auth,
724 "p".into(),
725 "global".into(),
726 Duration::from_secs(5),
727 Some(mock.uri()),
728 );
729 let c = provider
730 .generate(
731 "gemini-3-pro",
732 &req_with(VertexExt {
733 cached_content: Some("cachedContents/x".into()),
734 ..Default::default()
735 }),
736 None,
737 )
738 .await
739 .unwrap();
740 assert_eq!(c.content, "ok");
741 assert_eq!(c.input_tokens, 2);
742 }
743
744 #[tokio::test]
745 async fn stream_generate_yields_text_and_tool_done() {
746 use crate::routing::stream::{FinishReason, StreamItem};
747 use futures::StreamExt;
748 use wiremock::matchers::{method, path};
749 use wiremock::{Mock, MockServer, ResponseTemplate};
750
751 let sse = "data: {\"candidates\":[{\"content\":{\"role\":\"model\",\"parts\":[{\"text\":\"Hi\"}]}}]}\n\n\
752 data: {\"candidates\":[{\"content\":{\"role\":\"model\",\"parts\":[{\"functionCall\":{\"name\":\"f\",\"args\":{\"a\":1}}}]}}]}\n\n\
753 data: {\"candidates\":[{\"finishReason\":\"STOP\",\"content\":{\"role\":\"model\",\"parts\":[]}}],\"usageMetadata\":{\"promptTokenCount\":5,\"candidatesTokenCount\":4}}\n\n";
754 let mock = MockServer::start().await;
755 Mock::given(method("POST"))
756 .and(path("/v1/projects/p/locations/global/publishers/google/models/gemini-3-pro:streamGenerateContent"))
757 .respond_with(ResponseTemplate::new(200)
758 .insert_header("content-type", "text/event-stream")
759 .set_body_string(sse))
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(
767 auth,
768 "p".into(),
769 "global".into(),
770 Duration::from_secs(5),
771 Some(mock.uri()),
772 );
773
774 let req: ChatRequest = serde_json::from_value(serde_json::json!({
775 "model":"gemini-3-pro","messages":[{"role":"user","content":"hi"}],"stream":true}))
776 .unwrap();
777 let mut stream = std::pin::pin!(provider
778 .stream_generate("gemini-3-pro", &req, None)
779 .await
780 .expect("starts"));
781 let mut items = Vec::new();
782 while let Some(it) = stream.next().await {
783 items.push(it.unwrap());
784 }
785
786 assert!(items
787 .iter()
788 .any(|i| matches!(i, StreamItem::Delta(t) if t == "Hi")));
789 assert!(items
790 .iter()
791 .any(|i| matches!(i, StreamItem::ToolCallDelta { name: Some(n), .. } if n == "f")));
792 assert!(matches!(
793 items.last().unwrap(),
794 StreamItem::Done {
795 input_tokens: 5,
796 output_tokens: 4,
797 finish_reason: FinishReason::ToolCalls
798 }
799 ));
800 }
801
802 #[tokio::test]
803 async fn stream_generate_maps_vertex_4xx_to_bad_request() {
804 use wiremock::matchers::{method, path};
805 use wiremock::{Mock, MockServer, ResponseTemplate};
806
807 let mock = MockServer::start().await;
808 Mock::given(method("POST"))
809 .and(path("/v1/projects/p/locations/global/publishers/google/models/gemini-3-pro:streamGenerateContent"))
810 .respond_with(ResponseTemplate::new(400).set_body_string("bad responseSchema"))
811 .mount(&mock)
812 .await;
813
814 let auth = Arc::new(VertexAuth::with_fetcher(|| {
815 Box::pin(async { Ok(("test-token".into(), Duration::from_secs(3600))) })
816 }));
817 let provider = VertexNativeProvider::new(
818 auth,
819 "p".into(),
820 "global".into(),
821 Duration::from_secs(5),
822 Some(mock.uri()),
823 );
824 let req: ChatRequest = serde_json::from_value(serde_json::json!({
825 "model":"gemini-3-pro","messages":[{"role":"user","content":"hi"}],"stream":true}))
826 .unwrap();
827
828 let err = provider
829 .stream_generate("gemini-3-pro", &req, None)
830 .await
831 .err()
832 .expect("4xx should be an error");
833 assert!(
835 matches!(err, crate::error::GatewayError::BadRequest(_)),
836 "expected BadRequest, got {err:?}"
837 );
838 }
839}