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 parts = crate::routing::content_parts::content_to_vertex_parts(&m.content);
303 contents.push(json!({ "role": vrole, "parts": parts }));
304 }
305 }
306 }
307 if !media_parts.is_empty() {
309 if let Some(last) = contents.iter_mut().rev().find(|c| c["role"] == "user") {
310 if let Some(arr) = last["parts"].as_array_mut() {
311 arr.extend(media_parts);
312 }
313 } else {
314 contents.push(json!({ "role": "user", "parts": media_parts }));
315 }
316 }
317
318 let mut body = json!({ "contents": contents });
319
320 if let Some(cache) = &ext.cached_content {
321 body["cachedContent"] = json!(cache);
322 }
323 if let Some(schema) = &ext.response_schema {
324 body["generationConfig"] = json!({
325 "responseMimeType": "application/json",
326 "responseSchema": schema,
327 });
328 }
329 if let Some(t) = req.temperature {
330 body["generationConfig"]["temperature"] = json!(t);
331 }
332 if let Some(tools) = &req.tools {
333 let decls = tools.iter().filter_map(|t| {
334 let f = t.get("function")?;
335 Some(json!({
336 "name": f.get("name")?.as_str()?,
337 "description": f.get("description").and_then(|d| d.as_str()).unwrap_or(""),
338 "parameters": f.get("parameters").cloned().unwrap_or_else(|| json!({"type":"object"})),
339 }))
340 }).collect::<Vec<_>>();
341 if !decls.is_empty() {
342 body["tools"] = json!([{ "functionDeclarations": decls }]);
343 }
344 }
345 if let Some(choice) = &req.tool_choice {
346 let mode = match choice {
347 Value::String(s) if s == "none" => "NONE",
348 Value::String(s) if s == "required" => "ANY",
349 Value::String(_) => "AUTO",
350 Value::Object(_) => "ANY",
351 _ => "AUTO",
352 };
353 body["toolConfig"] = json!({ "functionCallingConfig": { "mode": mode } });
354 }
355
356 body
357}
358
359fn parse_response(
361 provider: &str,
362 model: &str,
363 v: &Value,
364) -> Result<Completion, crate::error::GatewayError> {
365 let content = v["candidates"][0]["content"]["parts"]
366 .as_array()
367 .map(|parts| {
368 parts
369 .iter()
370 .filter_map(|p| p["text"].as_str())
371 .collect::<Vec<_>>()
372 .join("")
373 })
374 .unwrap_or_default();
375 let usage = &v["usageMetadata"];
376 Ok(Completion {
377 provider: provider.to_string(),
378 model: model.to_string(),
379 content,
380 tool_calls: Vec::new(),
381 finish_reason: FinishReason::Stop,
382 input_tokens: usage["promptTokenCount"].as_u64().unwrap_or(0),
383 output_tokens: usage["candidatesTokenCount"].as_u64().unwrap_or(0),
384 })
385}
386
387fn map_vertex_finish(s: &str) -> FinishReason {
389 match s {
390 "MAX_TOKENS" => FinishReason::Length,
391 "STOP" => FinishReason::Stop,
392 _ => FinishReason::Stop,
393 }
394}
395
396pub fn vertex_chunk_to_items(chunk: &Value, tool_index: &mut u32) -> Vec<StreamItem> {
400 let mut out = Vec::new();
401 if let Some(parts) = chunk["candidates"][0]["content"]["parts"].as_array() {
402 for p in parts {
403 if let Some(text) = p["text"].as_str() {
404 if !text.is_empty() {
405 out.push(StreamItem::Delta(text.to_string()));
406 }
407 } else if let Some(fc) = p.get("functionCall") {
408 let name = fc["name"].as_str().unwrap_or_default().to_string();
409 let args = fc.get("args").cloned().unwrap_or_else(|| json!({}));
410 let i = *tool_index;
411 *tool_index += 1;
412 out.push(StreamItem::ToolCallDelta {
413 index: i,
414 id: Some(format!("call_{i}")),
415 name: Some(name),
416 args_fragment: args.to_string(),
417 });
418 }
419 }
420 }
421 if let Some(finish) = chunk["candidates"][0]["finishReason"].as_str() {
429 let usage = &chunk["usageMetadata"];
430 out.push(StreamItem::Done {
431 input_tokens: usage["promptTokenCount"].as_u64().unwrap_or(0),
432 output_tokens: usage["candidatesTokenCount"].as_u64().unwrap_or(0),
433 finish_reason: map_vertex_finish(finish),
434 });
435 }
436 out
437}
438
439#[cfg(test)]
440mod tests {
441 use super::*;
442
443 fn req_with(ext: VertexExt) -> ChatRequest {
444 let mut r: ChatRequest = serde_json::from_value(serde_json::json!({
445 "model": "gemini-pro", "messages": [{"role": "user", "content": "describe"}]
446 }))
447 .unwrap();
448 r.vertex = Some(ext);
449 r
450 }
451
452 #[test]
453 fn payload_includes_cached_content_and_schema_and_media() {
454 let ext = VertexExt {
455 cached_content: Some("cachedContents/abc".into()),
456 media_uris: Some(vec!["gs://bucket/v.mp4".into()]),
457 response_schema: Some(serde_json::json!({"type": "object"})),
458 };
459 let body = build_payload(&req_with(ext.clone()), &ext);
460 assert_eq!(
461 body["cachedContent"],
462 serde_json::json!("cachedContents/abc")
463 );
464 assert_eq!(
465 body["generationConfig"]["responseSchema"],
466 serde_json::json!({"type": "object"})
467 );
468 let parts = body["contents"][0]["parts"].as_array().unwrap();
469 assert!(parts
470 .iter()
471 .any(|p| p["fileData"]["fileUri"] == "gs://bucket/v.mp4"));
472 }
473
474 #[test]
475 fn parses_usage_from_vertex_response() {
476 let v = serde_json::json!({
477 "candidates": [{"content": {"parts": [{"text": "a"}, {"text": "b"}], "role": "model"}}],
478 "usageMetadata": {"promptTokenCount": 10, "candidatesTokenCount": 4}
479 });
480 let c = parse_response("vertex", "gemini-3-pro", &v).unwrap();
481 assert_eq!(c.content, "ab");
482 assert_eq!(c.input_tokens, 10);
483 assert_eq!(c.output_tokens, 4);
484 }
485
486 #[test]
487 fn payload_includes_tools_and_function_messages() {
488 let r: ChatRequest = serde_json::from_value(serde_json::json!({
489 "model": "gemini-pro",
490 "messages": [
491 {"role": "user", "content": "weather?"},
492 {"role": "assistant", "content": null, "tool_calls": [
493 {"id": "call_0", "type": "function", "function": {"name": "get_weather", "arguments": "{\"c\":\"SF\"}"}}]},
494 {"role": "tool", "tool_call_id": "call_0", "content": "21C"}
495 ],
496 "tools": [{"type": "function", "function": {"name": "get_weather",
497 "description": "Lookup", "parameters": {"type": "object"}}}],
498 "tool_choice": "auto"
499 })).unwrap();
500 let body = build_payload(&r, &VertexExt::default());
501 assert_eq!(
502 body["tools"][0]["functionDeclarations"][0]["name"],
503 "get_weather"
504 );
505 assert_eq!(body["toolConfig"]["functionCallingConfig"]["mode"], "AUTO");
506 let contents = body["contents"].as_array().unwrap();
507 assert!(contents
508 .iter()
509 .any(|c| c["parts"][0]["functionCall"]["name"] == "get_weather"));
510 assert!(contents
511 .iter()
512 .any(|c| c["parts"][0]["functionResponse"]["name"].is_string()));
513 }
514
515 #[test]
516 fn parses_vertex_chunk_text_and_functioncall() {
517 use crate::routing::stream::{FinishReason, StreamItem};
518 let mut idx = 0u32;
519
520 let text_chunk = serde_json::json!({
521 "candidates": [{"content": {"role": "model", "parts": [{"text": "Hi"}]}}]
522 });
523 let items = vertex_chunk_to_items(&text_chunk, &mut idx);
524 assert_eq!(items, vec![StreamItem::Delta("Hi".into())]);
525
526 let fc_chunk = serde_json::json!({
527 "candidates": [{"content": {"role": "model", "parts": [
528 {"functionCall": {"name": "get_weather", "args": {"c": "SF"}}}]}}]
529 });
530 let items = vertex_chunk_to_items(&fc_chunk, &mut idx);
531 assert_eq!(
532 items,
533 vec![StreamItem::ToolCallDelta {
534 index: 0,
535 id: Some("call_0".into()),
536 name: Some("get_weather".into()),
537 args_fragment: "{\"c\":\"SF\"}".into(),
538 }]
539 );
540
541 let final_chunk = serde_json::json!({
542 "candidates": [{"finishReason": "STOP", "content": {"role": "model", "parts": []}}],
543 "usageMetadata": {"promptTokenCount": 7, "candidatesTokenCount": 3}
544 });
545 let items = vertex_chunk_to_items(&final_chunk, &mut idx);
546 assert_eq!(
547 items,
548 vec![StreamItem::Done {
549 input_tokens: 7,
550 output_tokens: 3,
551 finish_reason: FinishReason::Stop
552 }]
553 );
554 }
555
556 #[tokio::test]
557 async fn generate_posts_to_vertex_url_with_bearer_and_parses() {
558 use wiremock::matchers::{header, method, path};
559 use wiremock::{Mock, MockServer, ResponseTemplate};
560
561 let mock = MockServer::start().await;
562 Mock::given(method("POST"))
563 .and(path("/v1/projects/p/locations/global/publishers/google/models/gemini-3-pro:generateContent"))
564 .and(header("authorization", "Bearer test-token"))
565 .respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({
566 "candidates": [{"content": {"parts": [{"text": "ok"}], "role": "model"}}],
567 "usageMetadata": {"promptTokenCount": 2, "candidatesTokenCount": 1}
568 })))
569 .mount(&mock)
570 .await;
571
572 let auth = Arc::new(VertexAuth::with_fetcher(|| {
573 Box::pin(async { Ok(("test-token".into(), Duration::from_secs(3600))) })
574 }));
575 let provider = VertexNativeProvider::new(
576 auth,
577 "p".into(),
578 "global".into(),
579 Duration::from_secs(5),
580 Some(mock.uri()),
581 );
582 let c = provider
583 .generate(
584 "gemini-3-pro",
585 &req_with(VertexExt {
586 cached_content: Some("cachedContents/x".into()),
587 ..Default::default()
588 }),
589 )
590 .await
591 .unwrap();
592 assert_eq!(c.content, "ok");
593 assert_eq!(c.input_tokens, 2);
594 }
595
596 #[tokio::test]
597 async fn stream_generate_yields_text_and_tool_done() {
598 use crate::routing::stream::{FinishReason, StreamItem};
599 use futures::StreamExt;
600 use wiremock::matchers::{method, path};
601 use wiremock::{Mock, MockServer, ResponseTemplate};
602
603 let sse = "data: {\"candidates\":[{\"content\":{\"role\":\"model\",\"parts\":[{\"text\":\"Hi\"}]}}]}\n\n\
604 data: {\"candidates\":[{\"content\":{\"role\":\"model\",\"parts\":[{\"functionCall\":{\"name\":\"f\",\"args\":{\"a\":1}}}]}}]}\n\n\
605 data: {\"candidates\":[{\"finishReason\":\"STOP\",\"content\":{\"role\":\"model\",\"parts\":[]}}],\"usageMetadata\":{\"promptTokenCount\":5,\"candidatesTokenCount\":4}}\n\n";
606 let mock = MockServer::start().await;
607 Mock::given(method("POST"))
608 .and(path("/v1/projects/p/locations/global/publishers/google/models/gemini-3-pro:streamGenerateContent"))
609 .respond_with(ResponseTemplate::new(200)
610 .insert_header("content-type", "text/event-stream")
611 .set_body_string(sse))
612 .mount(&mock)
613 .await;
614
615 let auth = Arc::new(VertexAuth::with_fetcher(|| {
616 Box::pin(async { Ok(("test-token".into(), Duration::from_secs(3600))) })
617 }));
618 let provider = VertexNativeProvider::new(
619 auth,
620 "p".into(),
621 "global".into(),
622 Duration::from_secs(5),
623 Some(mock.uri()),
624 );
625
626 let req: ChatRequest = serde_json::from_value(serde_json::json!({
627 "model":"gemini-3-pro","messages":[{"role":"user","content":"hi"}],"stream":true}))
628 .unwrap();
629 let mut stream = std::pin::pin!(provider
630 .stream_generate("gemini-3-pro", &req)
631 .await
632 .expect("starts"));
633 let mut items = Vec::new();
634 while let Some(it) = stream.next().await {
635 items.push(it.unwrap());
636 }
637
638 assert!(items
639 .iter()
640 .any(|i| matches!(i, StreamItem::Delta(t) if t == "Hi")));
641 assert!(items
642 .iter()
643 .any(|i| matches!(i, StreamItem::ToolCallDelta { name: Some(n), .. } if n == "f")));
644 assert!(matches!(
645 items.last().unwrap(),
646 StreamItem::Done {
647 input_tokens: 5,
648 output_tokens: 4,
649 finish_reason: FinishReason::ToolCalls
650 }
651 ));
652 }
653
654 #[tokio::test]
655 async fn stream_generate_maps_vertex_4xx_to_bad_request() {
656 use wiremock::matchers::{method, path};
657 use wiremock::{Mock, MockServer, ResponseTemplate};
658
659 let mock = MockServer::start().await;
660 Mock::given(method("POST"))
661 .and(path("/v1/projects/p/locations/global/publishers/google/models/gemini-3-pro:streamGenerateContent"))
662 .respond_with(ResponseTemplate::new(400).set_body_string("bad responseSchema"))
663 .mount(&mock)
664 .await;
665
666 let auth = Arc::new(VertexAuth::with_fetcher(|| {
667 Box::pin(async { Ok(("test-token".into(), Duration::from_secs(3600))) })
668 }));
669 let provider = VertexNativeProvider::new(
670 auth,
671 "p".into(),
672 "global".into(),
673 Duration::from_secs(5),
674 Some(mock.uri()),
675 );
676 let req: ChatRequest = serde_json::from_value(serde_json::json!({
677 "model":"gemini-3-pro","messages":[{"role":"user","content":"hi"}],"stream":true}))
678 .unwrap();
679
680 let err = provider
681 .stream_generate("gemini-3-pro", &req)
682 .await
683 .err()
684 .expect("4xx should be an error");
685 assert!(
687 matches!(err, crate::error::GatewayError::BadRequest(_)),
688 "expected BadRequest, got {err:?}"
689 );
690 }
691}