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,
20 endpoint_base: String,
22}
23
24impl VertexNativeProvider {
25 pub fn new(
26 auth: Arc<VertexAuth>,
27 project: String,
28 region: String,
29 request_timeout: Duration,
30 endpoint_override: Option<String>,
31 ) -> Self {
32 let endpoint_base = endpoint_override.unwrap_or_else(|| {
33 if region == "global" {
34 "https://aiplatform.googleapis.com".into()
35 } else {
36 format!("https://{region}-aiplatform.googleapis.com")
37 }
38 });
39 Self {
40 http: reqwest::Client::builder()
41 .timeout(request_timeout)
42 .build()
43 .unwrap(),
44 auth,
45 project,
46 region,
47 endpoint_base,
48 }
49 }
50
51 fn generate_url(&self, model: &str) -> String {
52 format!(
53 "{}/v1/projects/{}/locations/{}/publishers/google/models/{}:generateContent",
54 self.endpoint_base, self.project, self.region, model
55 )
56 }
57
58 pub async fn generate(
64 &self,
65 model: &str,
66 req: &ChatRequest,
67 ) -> Result<Completion, crate::error::GatewayError> {
68 let ext = req.vertex.clone().unwrap_or_default();
69 let payload = build_payload(req, &ext);
70 let token = self
71 .auth
72 .token()
73 .await
74 .map_err(|e| crate::error::GatewayError::Upstream {
75 status: 401,
76 body: format!("vertex auth: {e}"),
77 })?;
78
79 let resp = self
80 .http
81 .post(self.generate_url(model))
82 .bearer_auth(token)
83 .json(&payload)
84 .send()
85 .await
86 .map_err(|e| crate::error::GatewayError::Upstream {
87 status: 502,
88 body: e.to_string(),
89 })?;
90
91 let status = resp.status();
92 let value: Value = resp
93 .json()
94 .await
95 .map_err(|e| crate::error::GatewayError::Upstream {
96 status: status.as_u16(),
97 body: e.to_string(),
98 })?;
99 if !status.is_success() {
100 if status.is_client_error() {
101 return Err(crate::error::GatewayError::BadRequest(format!(
102 "vertex {}: {}",
103 status.as_u16(),
104 value
105 )));
106 }
107 return Err(crate::error::GatewayError::Upstream {
108 status: status.as_u16(),
109 body: value.to_string(),
110 });
111 }
112 parse_response("vertex", model, &value)
113 }
114
115 fn stream_url(&self, model: &str) -> String {
116 format!(
117 "{}/v1/projects/{}/locations/{}/publishers/google/models/{}:streamGenerateContent?alt=sse",
118 self.endpoint_base, self.project, self.region, model
119 )
120 }
121
122 pub async fn stream_generate(
129 &self,
130 model: &str,
131 req: &ChatRequest,
132 ) -> Result<
133 impl futures::Stream<Item = Result<StreamItem, crate::routing::executor::LegError>>,
134 crate::error::GatewayError,
135 > {
136 use crate::error::GatewayError;
137 use crate::routing::executor::LegError;
138 use crate::routing::stream::FinishReason;
139 use futures::StreamExt;
140
141 let ext = req.vertex.clone().unwrap_or_default();
142 let payload = build_payload(req, &ext);
143 let token = self
144 .auth
145 .token()
146 .await
147 .map_err(|e| GatewayError::Upstream {
148 status: 401,
149 body: format!("vertex auth: {e}"),
150 })?;
151
152 let resp = self
153 .http
154 .post(self.stream_url(model))
155 .bearer_auth(token)
156 .json(&payload)
157 .send()
158 .await
159 .map_err(|e| GatewayError::Upstream {
160 status: 502,
161 body: e.to_string(),
162 })?;
163
164 let status = resp.status();
165 if !status.is_success() {
166 let body = resp.text().await.unwrap_or_default();
167 if status.is_client_error() {
168 return Err(GatewayError::BadRequest(format!(
169 "vertex {}: {body}",
170 status.as_u16()
171 )));
172 }
173 return Err(GatewayError::Upstream {
174 status: status.as_u16(),
175 body,
176 });
177 }
178
179 struct St<S> {
183 inner: S,
184 buf: Vec<u8>,
189 tool_index: u32,
190 saw_tool: bool,
191 pending: std::collections::VecDeque<Result<StreamItem, LegError>>,
192 }
193 let state = St {
194 inner: Box::pin(resp.bytes_stream()),
195 buf: Vec::new(),
196 tool_index: 0,
197 saw_tool: false,
198 pending: std::collections::VecDeque::new(),
199 };
200
201 let items = futures::stream::unfold(state, |mut st| async move {
202 loop {
203 if let Some(item) = st.pending.pop_front() {
204 return Some((item, st));
205 }
206 match st.inner.next().await {
207 None => return None,
208 Some(Err(e)) => return Some((Err(LegError::MidStream(e.to_string())), st)),
209 Some(Ok(bytes)) => {
210 st.buf.extend_from_slice(&bytes);
211 while let Some(pos) = st.buf.windows(2).position(|w| w == b"\n\n") {
215 let event_bytes: Vec<u8> = st.buf.drain(..pos + 2).collect();
216 let event = String::from_utf8_lossy(&event_bytes);
217 for line in event.lines() {
218 let data = match line.strip_prefix("data:") {
219 Some(d) => d.trim(),
220 None => continue,
221 };
222 if data == "[DONE]" || data.is_empty() {
223 continue;
224 }
225 match serde_json::from_str::<serde_json::Value>(data) {
226 Ok(json) => {
227 for mut item in
228 vertex_chunk_to_items(&json, &mut st.tool_index)
229 {
230 if matches!(item, StreamItem::ToolCallDelta { .. }) {
231 st.saw_tool = true;
232 }
233 if let StreamItem::Done { finish_reason, .. } =
234 &mut item
235 {
236 if st.saw_tool {
237 *finish_reason = FinishReason::ToolCalls;
238 }
239 }
240 st.pending.push_back(Ok(item));
241 }
242 }
243 Err(e) => st.pending.push_back(Err(LegError::MidStream(
244 format!("bad sse json: {e}"),
245 ))),
246 }
247 }
248 }
249 }
251 }
252 }
253 });
254
255 Ok(items)
256 }
257}
258
259fn build_payload(req: &ChatRequest, ext: &VertexExt) -> Value {
262 let media_parts = ext
263 .media_uris
264 .iter()
265 .flatten()
266 .map(|uri| json!({ "fileData": { "fileUri": uri, "mimeType": "video/mp4" } }))
267 .collect::<Vec<_>>();
268
269 let mut contents: Vec<Value> = Vec::new();
271 for m in &req.messages {
272 match m.role.as_str() {
273 "assistant" if m.tool_calls.is_some() => {
274 let parts = m
275 .tool_calls
276 .as_ref()
277 .unwrap()
278 .iter()
279 .filter_map(|tc| {
280 let f = tc.get("function")?;
281 let name = f.get("name")?.as_str()?;
282 let raw = f.get("arguments").and_then(|a| a.as_str()).unwrap_or("{}");
283 let args: Value = serde_json::from_str(raw).unwrap_or_else(|_| json!({}));
284 Some(json!({ "functionCall": { "name": name, "args": args } }))
285 })
286 .collect::<Vec<_>>();
287 contents.push(json!({ "role": "model", "parts": parts }));
288 }
289 "tool" => {
290 let name = m.name.clone().unwrap_or_default();
291 let response: Value = m
292 .content
293 .as_str()
294 .map(|s| json!({ "content": s }))
295 .unwrap_or_else(|| json!({ "content": m.content.to_string() }));
296 contents.push(json!({ "role": "user", "parts": [
297 { "functionResponse": { "name": name, "response": response } }
298 ]}));
299 }
300 role => {
301 let vrole = if role == "assistant" { "model" } else { "user" };
302 let text = m
303 .content
304 .as_str()
305 .map(str::to_string)
306 .unwrap_or_else(|| m.content.to_string());
307 contents.push(json!({ "role": vrole, "parts": [{ "text": text }] }));
308 }
309 }
310 }
311 if !media_parts.is_empty() {
313 if let Some(last) = contents.iter_mut().rev().find(|c| c["role"] == "user") {
314 if let Some(arr) = last["parts"].as_array_mut() {
315 arr.extend(media_parts);
316 }
317 } else {
318 contents.push(json!({ "role": "user", "parts": media_parts }));
319 }
320 }
321
322 let mut body = json!({ "contents": contents });
323
324 if let Some(cache) = &ext.cached_content {
325 body["cachedContent"] = json!(cache);
326 }
327 if let Some(schema) = &ext.response_schema {
328 body["generationConfig"] = json!({
329 "responseMimeType": "application/json",
330 "responseSchema": schema,
331 });
332 }
333 if let Some(t) = req.temperature {
334 body["generationConfig"]["temperature"] = json!(t);
335 }
336 if let Some(tools) = &req.tools {
337 let decls = tools.iter().filter_map(|t| {
338 let f = t.get("function")?;
339 Some(json!({
340 "name": f.get("name")?.as_str()?,
341 "description": f.get("description").and_then(|d| d.as_str()).unwrap_or(""),
342 "parameters": f.get("parameters").cloned().unwrap_or_else(|| json!({"type":"object"})),
343 }))
344 }).collect::<Vec<_>>();
345 if !decls.is_empty() {
346 body["tools"] = json!([{ "functionDeclarations": decls }]);
347 }
348 }
349 if let Some(choice) = &req.tool_choice {
350 let mode = match choice {
351 Value::String(s) if s == "none" => "NONE",
352 Value::String(s) if s == "required" => "ANY",
353 Value::String(_) => "AUTO",
354 Value::Object(_) => "ANY",
355 _ => "AUTO",
356 };
357 body["toolConfig"] = json!({ "functionCallingConfig": { "mode": mode } });
358 }
359
360 body
361}
362
363fn parse_response(
365 provider: &str,
366 model: &str,
367 v: &Value,
368) -> Result<Completion, crate::error::GatewayError> {
369 let content = v["candidates"][0]["content"]["parts"]
370 .as_array()
371 .map(|parts| {
372 parts
373 .iter()
374 .filter_map(|p| p["text"].as_str())
375 .collect::<Vec<_>>()
376 .join("")
377 })
378 .unwrap_or_default();
379 let usage = &v["usageMetadata"];
380 Ok(Completion {
381 provider: provider.to_string(),
382 model: model.to_string(),
383 content,
384 tool_calls: Vec::new(),
385 finish_reason: FinishReason::Stop,
386 input_tokens: usage["promptTokenCount"].as_u64().unwrap_or(0),
387 output_tokens: usage["candidatesTokenCount"].as_u64().unwrap_or(0),
388 })
389}
390
391fn map_vertex_finish(s: &str) -> FinishReason {
393 match s {
394 "MAX_TOKENS" => FinishReason::Length,
395 "STOP" => FinishReason::Stop,
396 _ => FinishReason::Stop,
397 }
398}
399
400pub fn vertex_chunk_to_items(chunk: &Value, tool_index: &mut u32) -> Vec<StreamItem> {
404 let mut out = Vec::new();
405 if let Some(parts) = chunk["candidates"][0]["content"]["parts"].as_array() {
406 for p in parts {
407 if let Some(text) = p["text"].as_str() {
408 if !text.is_empty() {
409 out.push(StreamItem::Delta(text.to_string()));
410 }
411 } else if let Some(fc) = p.get("functionCall") {
412 let name = fc["name"].as_str().unwrap_or_default().to_string();
413 let args = fc.get("args").cloned().unwrap_or_else(|| json!({}));
414 let i = *tool_index;
415 *tool_index += 1;
416 out.push(StreamItem::ToolCallDelta {
417 index: i,
418 id: Some(format!("call_{i}")),
419 name: Some(name),
420 args_fragment: args.to_string(),
421 });
422 }
423 }
424 }
425 if let Some(finish) = chunk["candidates"][0]["finishReason"].as_str() {
433 let usage = &chunk["usageMetadata"];
434 out.push(StreamItem::Done {
435 input_tokens: usage["promptTokenCount"].as_u64().unwrap_or(0),
436 output_tokens: usage["candidatesTokenCount"].as_u64().unwrap_or(0),
437 finish_reason: map_vertex_finish(finish),
438 });
439 }
440 out
441}
442
443#[cfg(test)]
444mod tests {
445 use super::*;
446
447 fn req_with(ext: VertexExt) -> ChatRequest {
448 let mut r: ChatRequest = serde_json::from_value(serde_json::json!({
449 "model": "gemini-pro", "messages": [{"role": "user", "content": "describe"}]
450 }))
451 .unwrap();
452 r.vertex = Some(ext);
453 r
454 }
455
456 #[test]
457 fn payload_includes_cached_content_and_schema_and_media() {
458 let ext = VertexExt {
459 cached_content: Some("cachedContents/abc".into()),
460 media_uris: Some(vec!["gs://bucket/v.mp4".into()]),
461 response_schema: Some(serde_json::json!({"type": "object"})),
462 };
463 let body = build_payload(&req_with(ext.clone()), &ext);
464 assert_eq!(
465 body["cachedContent"],
466 serde_json::json!("cachedContents/abc")
467 );
468 assert_eq!(
469 body["generationConfig"]["responseSchema"],
470 serde_json::json!({"type": "object"})
471 );
472 let parts = body["contents"][0]["parts"].as_array().unwrap();
473 assert!(parts
474 .iter()
475 .any(|p| p["fileData"]["fileUri"] == "gs://bucket/v.mp4"));
476 }
477
478 #[test]
479 fn parses_usage_from_vertex_response() {
480 let v = serde_json::json!({
481 "candidates": [{"content": {"parts": [{"text": "a"}, {"text": "b"}], "role": "model"}}],
482 "usageMetadata": {"promptTokenCount": 10, "candidatesTokenCount": 4}
483 });
484 let c = parse_response("vertex", "gemini-3-pro", &v).unwrap();
485 assert_eq!(c.content, "ab");
486 assert_eq!(c.input_tokens, 10);
487 assert_eq!(c.output_tokens, 4);
488 }
489
490 #[test]
491 fn payload_includes_tools_and_function_messages() {
492 let r: ChatRequest = serde_json::from_value(serde_json::json!({
493 "model": "gemini-pro",
494 "messages": [
495 {"role": "user", "content": "weather?"},
496 {"role": "assistant", "content": null, "tool_calls": [
497 {"id": "call_0", "type": "function", "function": {"name": "get_weather", "arguments": "{\"c\":\"SF\"}"}}]},
498 {"role": "tool", "tool_call_id": "call_0", "content": "21C"}
499 ],
500 "tools": [{"type": "function", "function": {"name": "get_weather",
501 "description": "Lookup", "parameters": {"type": "object"}}}],
502 "tool_choice": "auto"
503 })).unwrap();
504 let body = build_payload(&r, &VertexExt::default());
505 assert_eq!(
506 body["tools"][0]["functionDeclarations"][0]["name"],
507 "get_weather"
508 );
509 assert_eq!(body["toolConfig"]["functionCallingConfig"]["mode"], "AUTO");
510 let contents = body["contents"].as_array().unwrap();
511 assert!(contents
512 .iter()
513 .any(|c| c["parts"][0]["functionCall"]["name"] == "get_weather"));
514 assert!(contents
515 .iter()
516 .any(|c| c["parts"][0]["functionResponse"]["name"].is_string()));
517 }
518
519 #[test]
520 fn parses_vertex_chunk_text_and_functioncall() {
521 use crate::routing::stream::{FinishReason, StreamItem};
522 let mut idx = 0u32;
523
524 let text_chunk = serde_json::json!({
525 "candidates": [{"content": {"role": "model", "parts": [{"text": "Hi"}]}}]
526 });
527 let items = vertex_chunk_to_items(&text_chunk, &mut idx);
528 assert_eq!(items, vec![StreamItem::Delta("Hi".into())]);
529
530 let fc_chunk = serde_json::json!({
531 "candidates": [{"content": {"role": "model", "parts": [
532 {"functionCall": {"name": "get_weather", "args": {"c": "SF"}}}]}}]
533 });
534 let items = vertex_chunk_to_items(&fc_chunk, &mut idx);
535 assert_eq!(
536 items,
537 vec![StreamItem::ToolCallDelta {
538 index: 0,
539 id: Some("call_0".into()),
540 name: Some("get_weather".into()),
541 args_fragment: "{\"c\":\"SF\"}".into(),
542 }]
543 );
544
545 let final_chunk = serde_json::json!({
546 "candidates": [{"finishReason": "STOP", "content": {"role": "model", "parts": []}}],
547 "usageMetadata": {"promptTokenCount": 7, "candidatesTokenCount": 3}
548 });
549 let items = vertex_chunk_to_items(&final_chunk, &mut idx);
550 assert_eq!(
551 items,
552 vec![StreamItem::Done {
553 input_tokens: 7,
554 output_tokens: 3,
555 finish_reason: FinishReason::Stop
556 }]
557 );
558 }
559
560 #[tokio::test]
561 async fn generate_posts_to_vertex_url_with_bearer_and_parses() {
562 use wiremock::matchers::{header, method, path};
563 use wiremock::{Mock, MockServer, ResponseTemplate};
564
565 let mock = MockServer::start().await;
566 Mock::given(method("POST"))
567 .and(path("/v1/projects/p/locations/global/publishers/google/models/gemini-3-pro:generateContent"))
568 .and(header("authorization", "Bearer test-token"))
569 .respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({
570 "candidates": [{"content": {"parts": [{"text": "ok"}], "role": "model"}}],
571 "usageMetadata": {"promptTokenCount": 2, "candidatesTokenCount": 1}
572 })))
573 .mount(&mock)
574 .await;
575
576 let auth = Arc::new(VertexAuth::with_fetcher(|| {
577 Box::pin(async { Ok(("test-token".into(), Duration::from_secs(3600))) })
578 }));
579 let provider = VertexNativeProvider::new(
580 auth,
581 "p".into(),
582 "global".into(),
583 Duration::from_secs(5),
584 Some(mock.uri()),
585 );
586 let c = provider
587 .generate(
588 "gemini-3-pro",
589 &req_with(VertexExt {
590 cached_content: Some("cachedContents/x".into()),
591 ..Default::default()
592 }),
593 )
594 .await
595 .unwrap();
596 assert_eq!(c.content, "ok");
597 assert_eq!(c.input_tokens, 2);
598 }
599
600 #[tokio::test]
601 async fn stream_generate_yields_text_and_tool_done() {
602 use crate::routing::stream::{FinishReason, StreamItem};
603 use futures::StreamExt;
604 use wiremock::matchers::{method, path};
605 use wiremock::{Mock, MockServer, ResponseTemplate};
606
607 let sse = "data: {\"candidates\":[{\"content\":{\"role\":\"model\",\"parts\":[{\"text\":\"Hi\"}]}}]}\n\n\
608 data: {\"candidates\":[{\"content\":{\"role\":\"model\",\"parts\":[{\"functionCall\":{\"name\":\"f\",\"args\":{\"a\":1}}}]}}]}\n\n\
609 data: {\"candidates\":[{\"finishReason\":\"STOP\",\"content\":{\"role\":\"model\",\"parts\":[]}}],\"usageMetadata\":{\"promptTokenCount\":5,\"candidatesTokenCount\":4}}\n\n";
610 let mock = MockServer::start().await;
611 Mock::given(method("POST"))
612 .and(path("/v1/projects/p/locations/global/publishers/google/models/gemini-3-pro:streamGenerateContent"))
613 .respond_with(ResponseTemplate::new(200)
614 .insert_header("content-type", "text/event-stream")
615 .set_body_string(sse))
616 .mount(&mock)
617 .await;
618
619 let auth = Arc::new(VertexAuth::with_fetcher(|| {
620 Box::pin(async { Ok(("test-token".into(), Duration::from_secs(3600))) })
621 }));
622 let provider = VertexNativeProvider::new(
623 auth,
624 "p".into(),
625 "global".into(),
626 Duration::from_secs(5),
627 Some(mock.uri()),
628 );
629
630 let req: ChatRequest = serde_json::from_value(serde_json::json!({
631 "model":"gemini-3-pro","messages":[{"role":"user","content":"hi"}],"stream":true}))
632 .unwrap();
633 let mut stream = std::pin::pin!(provider
634 .stream_generate("gemini-3-pro", &req)
635 .await
636 .expect("starts"));
637 let mut items = Vec::new();
638 while let Some(it) = stream.next().await {
639 items.push(it.unwrap());
640 }
641
642 assert!(items
643 .iter()
644 .any(|i| matches!(i, StreamItem::Delta(t) if t == "Hi")));
645 assert!(items
646 .iter()
647 .any(|i| matches!(i, StreamItem::ToolCallDelta { name: Some(n), .. } if n == "f")));
648 assert!(matches!(
649 items.last().unwrap(),
650 StreamItem::Done {
651 input_tokens: 5,
652 output_tokens: 4,
653 finish_reason: FinishReason::ToolCalls
654 }
655 ));
656 }
657
658 #[tokio::test]
659 async fn stream_generate_maps_vertex_4xx_to_bad_request() {
660 use wiremock::matchers::{method, path};
661 use wiremock::{Mock, MockServer, ResponseTemplate};
662
663 let mock = MockServer::start().await;
664 Mock::given(method("POST"))
665 .and(path("/v1/projects/p/locations/global/publishers/google/models/gemini-3-pro:streamGenerateContent"))
666 .respond_with(ResponseTemplate::new(400).set_body_string("bad responseSchema"))
667 .mount(&mock)
668 .await;
669
670 let auth = Arc::new(VertexAuth::with_fetcher(|| {
671 Box::pin(async { Ok(("test-token".into(), Duration::from_secs(3600))) })
672 }));
673 let provider = VertexNativeProvider::new(
674 auth,
675 "p".into(),
676 "global".into(),
677 Duration::from_secs(5),
678 Some(mock.uri()),
679 );
680 let req: ChatRequest = serde_json::from_value(serde_json::json!({
681 "model":"gemini-3-pro","messages":[{"role":"user","content":"hi"}],"stream":true}))
682 .unwrap();
683
684 let err = provider
685 .stream_generate("gemini-3-pro", &req)
686 .await
687 .err()
688 .expect("4xx should be an error");
689 assert!(
691 matches!(err, crate::error::GatewayError::BadRequest(_)),
692 "expected BadRequest, got {err:?}"
693 );
694 }
695}