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
19fn vertex_http_error(status: reqwest::StatusCode, body: String) -> crate::error::GatewayError {
23 if status.is_server_error()
24 || status == reqwest::StatusCode::TOO_MANY_REQUESTS
25 || status == reqwest::StatusCode::REQUEST_TIMEOUT
26 {
27 crate::error::GatewayError::Upstream {
28 status: status.as_u16(),
29 body,
30 }
31 } else if status.is_client_error() {
32 crate::error::GatewayError::BadRequest(format!("vertex {}: {body}", status.as_u16()))
33 } else {
34 crate::error::GatewayError::Upstream {
35 status: status.as_u16(),
36 body,
37 }
38 }
39}
40
41#[derive(Debug, Clone)]
42pub struct VertexNativeProvider {
43 http: reqwest::Client,
44 auth: Arc<VertexAuth>,
45 project: String,
46 region: String,
48 endpoint_override: Option<String>,
52}
53
54impl VertexNativeProvider {
55 pub fn new(
56 auth: Arc<VertexAuth>,
57 project: String,
58 region: String,
59 request_timeout: Duration,
60 endpoint_override: Option<String>,
61 ) -> Self {
62 Self {
63 http: reqwest::Client::builder()
64 .timeout(request_timeout)
65 .build()
66 .unwrap(),
67 auth,
68 project,
69 region,
70 endpoint_override,
71 }
72 }
73
74 fn endpoint_for(&self, region: &str) -> String {
78 if let Some(base) = &self.endpoint_override {
79 base.clone()
80 } else {
81 vertex_endpoint_base(region)
82 }
83 }
84
85 fn generate_url(&self, model: &str, region: &str) -> String {
86 format!(
87 "{}/v1/projects/{}/locations/{}/publishers/google/models/{}:generateContent",
88 self.endpoint_for(region),
89 self.project,
90 region,
91 model
92 )
93 }
94
95 pub async fn generate(
101 &self,
102 model: &str,
103 req: &ChatRequest,
104 region: Option<&str>,
105 ) -> Result<Completion, crate::error::GatewayError> {
106 let region = region.unwrap_or(&self.region);
107 let ext = req.vertex.clone().unwrap_or_default();
108 let payload = build_payload(req, &ext);
109 let token = self
110 .auth
111 .token()
112 .await
113 .map_err(|e| crate::error::GatewayError::Upstream {
114 status: 401,
115 body: format!("vertex auth: {e}"),
116 })?;
117
118 let resp = self
119 .http
120 .post(self.generate_url(model, region))
121 .bearer_auth(token)
122 .json(&payload)
123 .send()
124 .await
125 .map_err(|e| crate::error::GatewayError::Upstream {
126 status: 502,
127 body: e.to_string(),
128 })?;
129
130 let status = resp.status();
131 let value: Value = resp
132 .json()
133 .await
134 .map_err(|e| crate::error::GatewayError::Upstream {
135 status: status.as_u16(),
136 body: e.to_string(),
137 })?;
138 if !status.is_success() {
139 return Err(vertex_http_error(status, value.to_string()));
140 }
141 parse_response("vertex", model, &value)
142 }
143
144 pub async fn passthrough_request(
149 &self,
150 model: &str,
151 action: &str,
152 alt_sse: bool,
153 body: Value,
154 region: Option<&str>,
155 ) -> Result<reqwest::Response, crate::error::GatewayError> {
156 let region = region.unwrap_or(&self.region);
157 let alt = if alt_sse { "?alt=sse" } else { "" };
158 let url = format!(
159 "{}/v1/projects/{}/locations/{}/publishers/google/models/{}:{}{}",
160 self.endpoint_for(region),
161 self.project,
162 region,
163 model,
164 action,
165 alt
166 );
167 let token = self
168 .auth
169 .token()
170 .await
171 .map_err(|e| crate::error::GatewayError::Upstream {
172 status: 401,
173 body: format!("vertex auth: {e}"),
174 })?;
175 let mut request = self.http.post(url).bearer_auth(token).json(&body);
176 if alt_sse {
177 request = request.timeout(PASSTHROUGH_STREAM_TIMEOUT);
182 }
183 request
184 .send()
185 .await
186 .map_err(|e| crate::error::GatewayError::Upstream {
187 status: 502,
188 body: e.to_string(),
189 })
190 }
191
192 fn stream_url(&self, model: &str, region: &str) -> String {
193 format!(
194 "{}/v1/projects/{}/locations/{}/publishers/google/models/{}:streamGenerateContent?alt=sse",
195 self.endpoint_for(region),
196 self.project,
197 region,
198 model
199 )
200 }
201
202 pub async fn stream_generate(
210 &self,
211 model: &str,
212 req: &ChatRequest,
213 region: Option<&str>,
214 ) -> Result<
215 impl futures::Stream<Item = Result<StreamItem, crate::routing::executor::LegError>>,
216 crate::error::GatewayError,
217 > {
218 use crate::error::GatewayError;
219 use crate::routing::executor::LegError;
220 use crate::routing::stream::FinishReason;
221 use futures::StreamExt;
222
223 let region = region.unwrap_or(&self.region);
224 let ext = req.vertex.clone().unwrap_or_default();
225 let payload = build_payload(req, &ext);
226 let token = self
227 .auth
228 .token()
229 .await
230 .map_err(|e| GatewayError::Upstream {
231 status: 401,
232 body: format!("vertex auth: {e}"),
233 })?;
234
235 let resp = self
236 .http
237 .post(self.stream_url(model, region))
238 .bearer_auth(token)
239 .json(&payload)
240 .send()
241 .await
242 .map_err(|e| GatewayError::Upstream {
243 status: 502,
244 body: e.to_string(),
245 })?;
246
247 let status = resp.status();
248 if !status.is_success() {
249 let body = resp.text().await.unwrap_or_default();
250 return Err(vertex_http_error(status, body));
251 }
252
253 struct St<S> {
257 inner: S,
258 buf: Vec<u8>,
263 tool_index: u32,
264 saw_tool: bool,
265 pending: std::collections::VecDeque<Result<StreamItem, LegError>>,
266 }
267 let state = St {
268 inner: Box::pin(resp.bytes_stream()),
269 buf: Vec::new(),
270 tool_index: 0,
271 saw_tool: false,
272 pending: std::collections::VecDeque::new(),
273 };
274
275 let items = futures::stream::unfold(state, |mut st| async move {
276 loop {
277 if let Some(item) = st.pending.pop_front() {
278 return Some((item, st));
279 }
280 match st.inner.next().await {
281 None => return None,
282 Some(Err(e)) => return Some((Err(LegError::MidStream(e.to_string())), st)),
283 Some(Ok(bytes)) => {
284 st.buf.extend_from_slice(&bytes);
285 while let Some((pos, sep_len)) = next_sse_boundary(&st.buf) {
291 let event_bytes: Vec<u8> = st.buf.drain(..pos + sep_len).collect();
292 let event = String::from_utf8_lossy(&event_bytes);
293 for line in event.lines() {
294 let data = match line.strip_prefix("data:") {
295 Some(d) => d.trim(),
296 None => continue,
297 };
298 if data == "[DONE]" || data.is_empty() {
299 continue;
300 }
301 match serde_json::from_str::<serde_json::Value>(data) {
302 Ok(json) => {
303 for mut item in
304 vertex_chunk_to_items(&json, &mut st.tool_index)
305 {
306 if matches!(item, StreamItem::ToolCallDelta { .. }) {
307 st.saw_tool = true;
308 }
309 if let StreamItem::Done { finish_reason, .. } =
310 &mut item
311 {
312 if st.saw_tool {
313 *finish_reason = FinishReason::ToolCalls;
314 }
315 }
316 st.pending.push_back(Ok(item));
317 }
318 }
319 Err(e) => st.pending.push_back(Err(LegError::MidStream(
320 format!("bad sse json: {e}"),
321 ))),
322 }
323 }
324 }
325 }
327 }
328 }
329 });
330
331 Ok(items)
332 }
333}
334
335fn next_sse_boundary(buf: &[u8]) -> Option<(usize, usize)> {
340 let lf = buf
341 .windows(2)
342 .position(|w| w == b"\n\n")
343 .map(|p| (p, 2usize));
344 let crlf = buf
345 .windows(4)
346 .position(|w| w == b"\r\n\r\n")
347 .map(|p| (p, 4usize));
348 match (lf, crlf) {
349 (Some(a), Some(b)) => Some(if a.0 <= b.0 { a } else { b }),
350 (Some(a), None) => Some(a),
351 (None, Some(b)) => Some(b),
352 (None, None) => None,
353 }
354}
355
356fn build_payload(req: &ChatRequest, ext: &VertexExt) -> Value {
359 let media_parts = ext
360 .media_uris
361 .iter()
362 .flatten()
363 .map(|uri| json!({ "fileData": { "fileUri": uri, "mimeType": "video/mp4" } }))
364 .collect::<Vec<_>>();
365
366 let mut contents: Vec<Value> = Vec::new();
368 for m in &req.messages {
369 match m.role.as_str() {
370 "assistant" if m.tool_calls.is_some() => {
371 let parts = m
372 .tool_calls
373 .as_ref()
374 .unwrap()
375 .iter()
376 .filter_map(|tc| {
377 let f = tc.get("function")?;
378 let name = f.get("name")?.as_str()?;
379 let raw = f.get("arguments").and_then(|a| a.as_str()).unwrap_or("{}");
380 let args: Value = serde_json::from_str(raw).unwrap_or_else(|_| json!({}));
381 Some(json!({ "functionCall": { "name": name, "args": args } }))
382 })
383 .collect::<Vec<_>>();
384 contents.push(json!({ "role": "model", "parts": parts }));
385 }
386 "tool" => {
387 let name = m.name.clone().unwrap_or_default();
388 let response: Value = m
389 .content
390 .as_str()
391 .map(|s| json!({ "content": s }))
392 .unwrap_or_else(|| json!({ "content": m.content.to_string() }));
393 contents.push(json!({ "role": "user", "parts": [
394 { "functionResponse": { "name": name, "response": response } }
395 ]}));
396 }
397 role => {
398 let vrole = if role == "assistant" { "model" } else { "user" };
399 let parts = crate::routing::content_parts::content_to_vertex_parts(&m.content);
400 contents.push(json!({ "role": vrole, "parts": parts }));
401 }
402 }
403 }
404 if !media_parts.is_empty() {
406 if let Some(last) = contents.iter_mut().rev().find(|c| c["role"] == "user") {
407 if let Some(arr) = last["parts"].as_array_mut() {
408 arr.extend(media_parts);
409 }
410 } else {
411 contents.push(json!({ "role": "user", "parts": media_parts }));
412 }
413 }
414
415 let mut body = json!({ "contents": contents });
416
417 if let Some(cache) = &ext.cached_content {
418 body["cachedContent"] = json!(cache);
419 }
420 let mut gen_cfg = serde_json::Map::new();
424 if let Some(schema) = &ext.response_schema {
425 gen_cfg.insert("responseMimeType".into(), json!("application/json"));
426 gen_cfg.insert("responseSchema".into(), schema.clone());
427 }
428 if let Some(t) = req.temperature {
429 gen_cfg.insert("temperature".into(), json!(t));
430 }
431 if let Some(max) = req.max_tokens {
432 gen_cfg.insert("maxOutputTokens".into(), json!(max));
433 }
434 if let Some(thinking) = &ext.thinking_config {
435 gen_cfg.insert("thinkingConfig".into(), thinking.clone());
436 }
437 if !gen_cfg.is_empty() {
438 body["generationConfig"] = Value::Object(gen_cfg);
439 }
440 if let Some(tools) = &req.tools {
441 let decls = tools.iter().filter_map(|t| {
442 let f = t.get("function")?;
443 Some(json!({
444 "name": f.get("name")?.as_str()?,
445 "description": f.get("description").and_then(|d| d.as_str()).unwrap_or(""),
446 "parameters": f.get("parameters").cloned().unwrap_or_else(|| json!({"type":"object"})),
447 }))
448 }).collect::<Vec<_>>();
449 if !decls.is_empty() {
450 body["tools"] = json!([{ "functionDeclarations": decls }]);
451 }
452 }
453 if let Some(choice) = &req.tool_choice {
454 let mode = match choice {
455 Value::String(s) if s == "none" => "NONE",
456 Value::String(s) if s == "required" => "ANY",
457 Value::String(_) => "AUTO",
458 Value::Object(_) => "ANY",
459 _ => "AUTO",
460 };
461 body["toolConfig"] = json!({ "functionCallingConfig": { "mode": mode } });
462 }
463
464 body
465}
466
467fn is_thought_part(p: &Value) -> bool {
472 p.get("thought").and_then(Value::as_bool).unwrap_or(false)
473}
474
475fn parse_response(
477 provider: &str,
478 model: &str,
479 v: &Value,
480) -> Result<Completion, crate::error::GatewayError> {
481 let content = v["candidates"][0]["content"]["parts"]
482 .as_array()
483 .map(|parts| {
484 parts
485 .iter()
486 .filter(|p| !is_thought_part(p))
487 .filter_map(|p| p["text"].as_str())
488 .collect::<Vec<_>>()
489 .join("")
490 })
491 .unwrap_or_default();
492 let usage = &v["usageMetadata"];
493 Ok(Completion {
494 provider: provider.to_string(),
495 model: model.to_string(),
496 content,
497 tool_calls: Vec::new(),
498 finish_reason: FinishReason::Stop,
499 input_tokens: usage["promptTokenCount"].as_u64().unwrap_or(0),
500 output_tokens: usage["candidatesTokenCount"].as_u64().unwrap_or(0),
501 })
502}
503
504fn map_vertex_finish(s: &str) -> FinishReason {
506 match s {
507 "MAX_TOKENS" => FinishReason::Length,
508 "STOP" => FinishReason::Stop,
509 _ => FinishReason::Stop,
510 }
511}
512
513pub fn vertex_chunk_to_items(chunk: &Value, tool_index: &mut u32) -> Vec<StreamItem> {
517 let mut out = Vec::new();
518 if let Some(parts) = chunk["candidates"][0]["content"]["parts"].as_array() {
519 for p in parts {
520 if is_thought_part(p) {
521 continue;
522 }
523 if let Some(text) = p["text"].as_str() {
524 if !text.is_empty() {
525 out.push(StreamItem::Delta(text.to_string()));
526 }
527 } else if let Some(fc) = p.get("functionCall") {
528 let name = fc["name"].as_str().unwrap_or_default().to_string();
529 let args = fc.get("args").cloned().unwrap_or_else(|| json!({}));
530 let i = *tool_index;
531 *tool_index += 1;
532 out.push(StreamItem::ToolCallDelta {
533 index: i,
534 id: Some(format!("call_{i}")),
535 name: Some(name),
536 args_fragment: args.to_string(),
537 });
538 }
539 }
540 }
541 if let Some(finish) = chunk["candidates"][0]["finishReason"].as_str() {
549 let usage = &chunk["usageMetadata"];
550 out.push(StreamItem::Done {
551 input_tokens: usage["promptTokenCount"].as_u64().unwrap_or(0),
552 output_tokens: usage["candidatesTokenCount"].as_u64().unwrap_or(0),
553 finish_reason: map_vertex_finish(finish),
554 });
555 }
556 out
557}
558
559#[cfg(test)]
560mod tests {
561 use super::*;
562
563 fn req_with(ext: VertexExt) -> ChatRequest {
564 let mut r: ChatRequest = serde_json::from_value(serde_json::json!({
565 "model": "gemini-pro", "messages": [{"role": "user", "content": "describe"}]
566 }))
567 .unwrap();
568 r.vertex = Some(ext);
569 r
570 }
571
572 #[test]
573 fn payload_includes_cached_content_and_schema_and_media() {
574 let ext = VertexExt {
575 cached_content: Some("cachedContents/abc".into()),
576 media_uris: Some(vec!["gs://bucket/v.mp4".into()]),
577 response_schema: Some(serde_json::json!({"type": "object"})),
578 ..Default::default()
579 };
580 let body = build_payload(&req_with(ext.clone()), &ext);
581 assert_eq!(
582 body["cachedContent"],
583 serde_json::json!("cachedContents/abc")
584 );
585 assert_eq!(
586 body["generationConfig"]["responseSchema"],
587 serde_json::json!({"type": "object"})
588 );
589 let parts = body["contents"][0]["parts"].as_array().unwrap();
590 assert!(parts
591 .iter()
592 .any(|p| p["fileData"]["fileUri"] == "gs://bucket/v.mp4"));
593 }
594
595 #[test]
596 fn payload_includes_max_output_tokens_and_thinking_config() {
597 let ext = VertexExt {
598 thinking_config: Some(serde_json::json!({ "thinkingLevel": "low" })),
599 ..Default::default()
600 };
601 let mut req = req_with(ext.clone());
602 req.max_tokens = Some(8192);
603 let body = build_payload(&req, &ext);
604 assert_eq!(
605 body["generationConfig"]["maxOutputTokens"],
606 serde_json::json!(8192)
607 );
608 assert_eq!(
609 body["generationConfig"]["thinkingConfig"],
610 serde_json::json!({ "thinkingLevel": "low" })
611 );
612 }
613
614 #[test]
615 fn vertex_chunk_skips_thought_parts() {
616 use crate::routing::stream::StreamItem;
617 let mut idx = 0u32;
618 let chunk = serde_json::json!({
619 "candidates": [{"content": {"role": "model", "parts": [
620 {"text": "internal reasoning", "thought": true},
621 {"text": "answer"}
622 ]}}]
623 });
624 let items = vertex_chunk_to_items(&chunk, &mut idx);
625 assert_eq!(items, vec![StreamItem::Delta("answer".into())]);
626 }
627
628 #[test]
629 fn parse_response_skips_thought_parts() {
630 let v = serde_json::json!({
631 "candidates": [{"content": {"parts": [
632 {"text": "reasoning", "thought": true},
633 {"text": "real"}
634 ], "role": "model"}}],
635 "usageMetadata": {"promptTokenCount": 1, "candidatesTokenCount": 1}
636 });
637 let c = parse_response("vertex", "gemini-3-pro", &v).unwrap();
638 assert_eq!(c.content, "real");
639 }
640
641 #[test]
642 fn parses_usage_from_vertex_response() {
643 let v = serde_json::json!({
644 "candidates": [{"content": {"parts": [{"text": "a"}, {"text": "b"}], "role": "model"}}],
645 "usageMetadata": {"promptTokenCount": 10, "candidatesTokenCount": 4}
646 });
647 let c = parse_response("vertex", "gemini-3-pro", &v).unwrap();
648 assert_eq!(c.content, "ab");
649 assert_eq!(c.input_tokens, 10);
650 assert_eq!(c.output_tokens, 4);
651 }
652
653 #[test]
654 fn payload_includes_tools_and_function_messages() {
655 let r: ChatRequest = serde_json::from_value(serde_json::json!({
656 "model": "gemini-pro",
657 "messages": [
658 {"role": "user", "content": "weather?"},
659 {"role": "assistant", "content": null, "tool_calls": [
660 {"id": "call_0", "type": "function", "function": {"name": "get_weather", "arguments": "{\"c\":\"SF\"}"}}]},
661 {"role": "tool", "tool_call_id": "call_0", "content": "21C"}
662 ],
663 "tools": [{"type": "function", "function": {"name": "get_weather",
664 "description": "Lookup", "parameters": {"type": "object"}}}],
665 "tool_choice": "auto"
666 })).unwrap();
667 let body = build_payload(&r, &VertexExt::default());
668 assert_eq!(
669 body["tools"][0]["functionDeclarations"][0]["name"],
670 "get_weather"
671 );
672 assert_eq!(body["toolConfig"]["functionCallingConfig"]["mode"], "AUTO");
673 let contents = body["contents"].as_array().unwrap();
674 assert!(contents
675 .iter()
676 .any(|c| c["parts"][0]["functionCall"]["name"] == "get_weather"));
677 assert!(contents
678 .iter()
679 .any(|c| c["parts"][0]["functionResponse"]["name"].is_string()));
680 }
681
682 #[test]
683 fn parses_vertex_chunk_text_and_functioncall() {
684 use crate::routing::stream::{FinishReason, StreamItem};
685 let mut idx = 0u32;
686
687 let text_chunk = serde_json::json!({
688 "candidates": [{"content": {"role": "model", "parts": [{"text": "Hi"}]}}]
689 });
690 let items = vertex_chunk_to_items(&text_chunk, &mut idx);
691 assert_eq!(items, vec![StreamItem::Delta("Hi".into())]);
692
693 let fc_chunk = serde_json::json!({
694 "candidates": [{"content": {"role": "model", "parts": [
695 {"functionCall": {"name": "get_weather", "args": {"c": "SF"}}}]}}]
696 });
697 let items = vertex_chunk_to_items(&fc_chunk, &mut idx);
698 assert_eq!(
699 items,
700 vec![StreamItem::ToolCallDelta {
701 index: 0,
702 id: Some("call_0".into()),
703 name: Some("get_weather".into()),
704 args_fragment: "{\"c\":\"SF\"}".into(),
705 }]
706 );
707
708 let final_chunk = serde_json::json!({
709 "candidates": [{"finishReason": "STOP", "content": {"role": "model", "parts": []}}],
710 "usageMetadata": {"promptTokenCount": 7, "candidatesTokenCount": 3}
711 });
712 let items = vertex_chunk_to_items(&final_chunk, &mut idx);
713 assert_eq!(
714 items,
715 vec![StreamItem::Done {
716 input_tokens: 7,
717 output_tokens: 3,
718 finish_reason: FinishReason::Stop
719 }]
720 );
721 }
722
723 #[test]
724 fn endpoint_for_region_picks_global_multi_or_single_host() {
725 let auth = Arc::new(VertexAuth::with_fetcher(|| {
726 Box::pin(async { Ok(("t".into(), Duration::from_secs(3600))) })
727 }));
728 let provider = VertexNativeProvider::new(
729 auth,
730 "p".into(),
731 "global".into(),
732 Duration::from_secs(5),
733 None,
734 );
735 assert_eq!(
736 provider.endpoint_for("global"),
737 "https://aiplatform.googleapis.com"
738 );
739 assert_eq!(
740 provider.endpoint_for("us"),
741 "https://aiplatform.us.rep.googleapis.com"
742 );
743 assert_eq!(
744 provider.endpoint_for("eu"),
745 "https://aiplatform.eu.rep.googleapis.com"
746 );
747 assert_eq!(
748 provider.endpoint_for("us-central1"),
749 "https://us-central1-aiplatform.googleapis.com"
750 );
751 }
752
753 #[tokio::test]
754 async fn generate_uses_per_leg_region_override_in_url() {
755 use wiremock::matchers::{method, path};
756 use wiremock::{Mock, MockServer, ResponseTemplate};
757
758 let mock = MockServer::start().await;
759 Mock::given(method("POST"))
760 .and(path("/v1/projects/p/locations/us-central1/publishers/google/models/gemini-x:generateContent"))
761 .respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({
762 "candidates": [{"content": {"parts": [{"text": "ok"}], "role": "model"}}],
763 "usageMetadata": {"promptTokenCount": 1, "candidatesTokenCount": 1}
764 })))
765 .mount(&mock)
766 .await;
767
768 let auth = Arc::new(VertexAuth::with_fetcher(|| {
769 Box::pin(async { Ok(("test-token".into(), Duration::from_secs(3600))) })
770 }));
771 let provider = VertexNativeProvider::new(
773 auth,
774 "p".into(),
775 "global".into(),
776 Duration::from_secs(5),
777 Some(mock.uri()),
778 );
779 let c = provider
780 .generate(
781 "gemini-x",
782 &req_with(VertexExt::default()),
783 Some("us-central1"),
784 )
785 .await
786 .unwrap();
787 assert_eq!(c.content, "ok");
788 }
789
790 #[tokio::test]
791 async fn generate_uses_multi_region_us_override_in_url() {
792 use wiremock::matchers::{method, path};
793 use wiremock::{Mock, MockServer, ResponseTemplate};
794
795 let mock = MockServer::start().await;
796 Mock::given(method("POST"))
797 .and(path(
798 "/v1/projects/p/locations/us/publishers/google/models/gemini-lite:generateContent",
799 ))
800 .respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({
801 "candidates": [{"content": {"parts": [{"text": "ok"}], "role": "model"}}],
802 "usageMetadata": {"promptTokenCount": 1, "candidatesTokenCount": 1}
803 })))
804 .mount(&mock)
805 .await;
806
807 let auth = Arc::new(VertexAuth::with_fetcher(|| {
808 Box::pin(async { Ok(("test-token".into(), Duration::from_secs(3600))) })
809 }));
810 let provider = VertexNativeProvider::new(
811 auth,
812 "p".into(),
813 "global".into(),
814 Duration::from_secs(5),
815 Some(mock.uri()),
816 );
817 let c = provider
818 .generate("gemini-lite", &req_with(VertexExt::default()), Some("us"))
819 .await
820 .unwrap();
821 assert_eq!(c.content, "ok");
822 }
823
824 #[tokio::test]
825 async fn generate_posts_to_vertex_url_with_bearer_and_parses() {
826 use wiremock::matchers::{header, method, path};
827 use wiremock::{Mock, MockServer, ResponseTemplate};
828
829 let mock = MockServer::start().await;
830 Mock::given(method("POST"))
831 .and(path("/v1/projects/p/locations/global/publishers/google/models/gemini-3-pro:generateContent"))
832 .and(header("authorization", "Bearer test-token"))
833 .respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({
834 "candidates": [{"content": {"parts": [{"text": "ok"}], "role": "model"}}],
835 "usageMetadata": {"promptTokenCount": 2, "candidatesTokenCount": 1}
836 })))
837 .mount(&mock)
838 .await;
839
840 let auth = Arc::new(VertexAuth::with_fetcher(|| {
841 Box::pin(async { Ok(("test-token".into(), Duration::from_secs(3600))) })
842 }));
843 let provider = VertexNativeProvider::new(
844 auth,
845 "p".into(),
846 "global".into(),
847 Duration::from_secs(5),
848 Some(mock.uri()),
849 );
850 let c = provider
851 .generate(
852 "gemini-3-pro",
853 &req_with(VertexExt {
854 cached_content: Some("cachedContents/x".into()),
855 ..Default::default()
856 }),
857 None,
858 )
859 .await
860 .unwrap();
861 assert_eq!(c.content, "ok");
862 assert_eq!(c.input_tokens, 2);
863 }
864
865 #[tokio::test]
866 async fn stream_generate_yields_text_and_tool_done() {
867 use crate::routing::stream::{FinishReason, StreamItem};
868 use futures::StreamExt;
869 use wiremock::matchers::{method, path};
870 use wiremock::{Mock, MockServer, ResponseTemplate};
871
872 let sse = "data: {\"candidates\":[{\"content\":{\"role\":\"model\",\"parts\":[{\"text\":\"Hi\"}]}}]}\n\n\
873 data: {\"candidates\":[{\"content\":{\"role\":\"model\",\"parts\":[{\"functionCall\":{\"name\":\"f\",\"args\":{\"a\":1}}}]}}]}\n\n\
874 data: {\"candidates\":[{\"finishReason\":\"STOP\",\"content\":{\"role\":\"model\",\"parts\":[]}}],\"usageMetadata\":{\"promptTokenCount\":5,\"candidatesTokenCount\":4}}\n\n";
875 let mock = MockServer::start().await;
876 Mock::given(method("POST"))
877 .and(path("/v1/projects/p/locations/global/publishers/google/models/gemini-3-pro:streamGenerateContent"))
878 .respond_with(ResponseTemplate::new(200)
879 .insert_header("content-type", "text/event-stream")
880 .set_body_string(sse))
881 .mount(&mock)
882 .await;
883
884 let auth = Arc::new(VertexAuth::with_fetcher(|| {
885 Box::pin(async { Ok(("test-token".into(), Duration::from_secs(3600))) })
886 }));
887 let provider = VertexNativeProvider::new(
888 auth,
889 "p".into(),
890 "global".into(),
891 Duration::from_secs(5),
892 Some(mock.uri()),
893 );
894
895 let req: ChatRequest = serde_json::from_value(serde_json::json!({
896 "model":"gemini-3-pro","messages":[{"role":"user","content":"hi"}],"stream":true}))
897 .unwrap();
898 let mut stream = std::pin::pin!(provider
899 .stream_generate("gemini-3-pro", &req, None)
900 .await
901 .expect("starts"));
902 let mut items = Vec::new();
903 while let Some(it) = stream.next().await {
904 items.push(it.unwrap());
905 }
906
907 assert!(items
908 .iter()
909 .any(|i| matches!(i, StreamItem::Delta(t) if t == "Hi")));
910 assert!(items
911 .iter()
912 .any(|i| matches!(i, StreamItem::ToolCallDelta { name: Some(n), .. } if n == "f")));
913 assert!(matches!(
914 items.last().unwrap(),
915 StreamItem::Done {
916 input_tokens: 5,
917 output_tokens: 4,
918 finish_reason: FinishReason::ToolCalls
919 }
920 ));
921 }
922
923 #[tokio::test]
924 async fn stream_generate_parses_crlf_terminated_events() {
925 use crate::routing::stream::{FinishReason, StreamItem};
932 use futures::StreamExt;
933 use wiremock::matchers::{method, path};
934 use wiremock::{Mock, MockServer, ResponseTemplate};
935
936 let sse = "data: {\"candidates\":[{\"content\":{\"role\":\"model\",\"parts\":[{\"text\":\"{\\\"answer\\\":\\\"hi\\\"}\"}]}}],\"usageMetadata\":{\"trafficType\":\"ON_DEMAND\"}}\r\n\r\n\
937 data: {\"candidates\":[{\"finishReason\":\"STOP\",\"content\":{\"role\":\"model\",\"parts\":[{\"text\":\"\",\"thoughtSignature\":\"abc\"}]}}],\"usageMetadata\":{\"promptTokenCount\":56,\"candidatesTokenCount\":8}}\r\n\r\n";
938 let mock = MockServer::start().await;
939 Mock::given(method("POST"))
940 .and(path("/v1/projects/p/locations/global/publishers/google/models/gemini-3-pro:streamGenerateContent"))
941 .respond_with(ResponseTemplate::new(200)
942 .insert_header("content-type", "text/event-stream")
943 .set_body_string(sse))
944 .mount(&mock)
945 .await;
946
947 let auth = Arc::new(VertexAuth::with_fetcher(|| {
948 Box::pin(async { Ok(("test-token".into(), Duration::from_secs(3600))) })
949 }));
950 let provider = VertexNativeProvider::new(
951 auth,
952 "p".into(),
953 "global".into(),
954 Duration::from_secs(5),
955 Some(mock.uri()),
956 );
957 let req: ChatRequest = serde_json::from_value(serde_json::json!({
958 "model":"gemini-3-pro","messages":[{"role":"user","content":"hi"}],"stream":true}))
959 .unwrap();
960 let mut stream = std::pin::pin!(provider
961 .stream_generate("gemini-3-pro", &req, None)
962 .await
963 .expect("starts"));
964 let mut items = Vec::new();
965 while let Some(it) = stream.next().await {
966 items.push(it.unwrap());
967 }
968
969 let text: String = items
970 .iter()
971 .filter_map(|i| match i {
972 StreamItem::Delta(t) => Some(t.clone()),
973 _ => None,
974 })
975 .collect();
976 assert_eq!(
977 text, "{\"answer\":\"hi\"}",
978 "answer text must survive CRLF events"
979 );
980 assert!(matches!(
981 items.last().unwrap(),
982 StreamItem::Done {
983 input_tokens: 56,
984 output_tokens: 8,
985 finish_reason: FinishReason::Stop
986 }
987 ));
988 }
989
990 #[tokio::test]
991 async fn stream_generate_maps_vertex_4xx_to_bad_request() {
992 use wiremock::matchers::{method, path};
993 use wiremock::{Mock, MockServer, ResponseTemplate};
994
995 let mock = MockServer::start().await;
996 Mock::given(method("POST"))
997 .and(path("/v1/projects/p/locations/global/publishers/google/models/gemini-3-pro:streamGenerateContent"))
998 .respond_with(ResponseTemplate::new(400).set_body_string("bad responseSchema"))
999 .mount(&mock)
1000 .await;
1001
1002 let auth = Arc::new(VertexAuth::with_fetcher(|| {
1003 Box::pin(async { Ok(("test-token".into(), Duration::from_secs(3600))) })
1004 }));
1005 let provider = VertexNativeProvider::new(
1006 auth,
1007 "p".into(),
1008 "global".into(),
1009 Duration::from_secs(5),
1010 Some(mock.uri()),
1011 );
1012 let req: ChatRequest = serde_json::from_value(serde_json::json!({
1013 "model":"gemini-3-pro","messages":[{"role":"user","content":"hi"}],"stream":true}))
1014 .unwrap();
1015
1016 let err = provider
1017 .stream_generate("gemini-3-pro", &req, None)
1018 .await
1019 .err()
1020 .expect("4xx should be an error");
1021 assert!(
1023 matches!(err, crate::error::GatewayError::BadRequest(_)),
1024 "expected BadRequest, got {err:?}"
1025 );
1026 }
1027}