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