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