From 494aab6d149c5673fab425bf672c00c069dd6e17 Mon Sep 17 00:00:00 2001 From: Andre Bandarra Date: Sun, 20 Sep 2026 16:42:55 +0100 Subject: [PATCH] feat: add SSE streaming, thinking models, and modernized Part struct --- Cargo.toml | 24 +- examples/generate_image.rs | 7 +- examples/google-search-retrieval.rs | 2 +- examples/google-search.rs | 2 +- examples/json-schema.rs | 2 +- examples/safety-setting.rs | 2 +- examples/text-from-text-streaming.rs | 6 +- src/client.rs | 527 +++++++++++---------------- src/dialogue.rs | 92 ++--- src/error.rs | 161 ++++---- src/lib.rs | 1 + src/network/event_source.rs | 108 ++++++ src/network/mod.rs | 1 + src/token_provider.rs | 36 +- src/types/common.rs | 177 +++++---- src/types/count_tokens.rs | 2 +- src/types/generate_content.rs | 126 ++++--- src/types/predict_image.rs | 123 +++---- 18 files changed, 703 insertions(+), 696 deletions(-) create mode 100644 src/network/event_source.rs create mode 100644 src/network/mod.rs diff --git a/Cargo.toml b/Cargo.toml index a2dcfc0..8fe851b 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -1,26 +1,26 @@ [package] name = "gemini-rs" version = "0.1.0" -edition = "2021" +edition = "2024" # See more keys and their definitions at https://doc.rust-lang.org/cargo/reference/manifest.html [dependencies] deadqueue = "0.2" gcp_auth = "0.12" -reqwest = { version = "0.12", features = ["json", "gzip"] } -reqwest-eventsource = "0.6" +reqwest = { version = "0.13", features = ["json", "gzip", "stream"] } +tokio-util = "0.7" serde = { version = "1", features = ["derive"] } -serde_json = { version = "1"} -serde_with = { version = "3.9", features = ["base64"]} +serde_json = { version = "1" } +serde_with = { version = "3.21", features = ["base64"] } tracing = "0.1" tokio = { version = "1" } -tokio-stream = "0.1.17" +tokio-stream = "0.1.18" [dev-dependencies] -console = "0.15.8" -dialoguer = "0.11.0" -image = "0.25.2" -indicatif = "0.17.8" -tokio = { version = "1.37.0", features = ["full"] } -tracing-subscriber = "0.3.18" +console = "0.16.4" +dialoguer = "0.12.0" +image = "0.25.10" +indicatif = "0.18.6" +tokio = { version = "1.52.3", features = ["full"] } +tracing-subscriber = "0.3.23" diff --git a/examples/generate_image.rs b/examples/generate_image.rs index f894f95..9a79d2d 100644 --- a/examples/generate_image.rs +++ b/examples/generate_image.rs @@ -1,9 +1,8 @@ use std::{error::Error, io::Cursor}; use gemini_rs::prelude::{ - GeminiClient, PersonGeneration, PredictImageRequest, PredictImageRequestParameters, + GeminiClient, PredictImageRequest, PredictImageRequestParameters, PredictImageRequestParametersOutputOptions, PredictImageRequestPrompt, - PredictImageSafetySetting, }; use image::{ImageFormat, ImageReader}; @@ -35,8 +34,8 @@ pub async fn main() -> Result<(), Box> { mime_type: Some("image/jpeg".to_string()), compression_quality: Some(75), }), - person_generation: Some(PersonGeneration::AllowAll), - safety_setting: Some(PredictImageSafetySetting::BlockOnlyHigh), + person_generation: Some("allow_all".to_string()), + safety_setting: Some("block_only_high".to_string()), ..Default::default() }, }; diff --git a/examples/google-search-retrieval.rs b/examples/google-search-retrieval.rs index f6832dc..b9b59e3 100644 --- a/examples/google-search-retrieval.rs +++ b/examples/google-search-retrieval.rs @@ -20,7 +20,7 @@ async fn main() -> Result<(), Box> { let request = GenerateContentRequest { contents: vec![Content { role: Some(Role::User), - parts: Some(vec![Part::Text(prompt.to_string())]), + parts: Some(vec![Part::text(prompt)]), }], tools: Some(vec![Tools { google_search_retrieval: Some(GoogleSearchRetrieval::default()), diff --git a/examples/google-search.rs b/examples/google-search.rs index fb38d69..1bcc45c 100644 --- a/examples/google-search.rs +++ b/examples/google-search.rs @@ -20,7 +20,7 @@ async fn main() -> Result<(), Box> { let request = GenerateContentRequest { contents: vec![Content { role: Some(Role::User), - parts: Some(vec![Part::Text(prompt.to_string())]), + parts: Some(vec![Part::text(prompt)]), }], tools: Some(vec![Tools { google_search: Some(GoogleSearch::default()), diff --git a/examples/json-schema.rs b/examples/json-schema.rs index c56cf23..5661307 100644 --- a/examples/json-schema.rs +++ b/examples/json-schema.rs @@ -19,7 +19,7 @@ async fn main() -> Result<(), Box> { let request = GenerateContentRequest { contents: vec![Content { role: Some(Role::User), - parts: Some(vec![Part::Text(prompt.to_string())]), + parts: Some(vec![Part::text(prompt)]), }], generation_config: Some(GenerationConfig { response_mime_type: Some("application/json".to_string()), diff --git a/examples/safety-setting.rs b/examples/safety-setting.rs index a4840d0..12d9a82 100644 --- a/examples/safety-setting.rs +++ b/examples/safety-setting.rs @@ -21,7 +21,7 @@ async fn main() -> Result<(), Box> { let request = GenerateContentRequest { contents: vec![Content { role: Some(Role::User), - parts: Some(vec![Part::Text(prompt.to_string())]), + parts: Some(vec![Part::text(prompt)]), }], safety_settings: Some(vec![SafetySetting { category: HarmCategory::HateSpeech, diff --git a/examples/text-from-text-streaming.rs b/examples/text-from-text-streaming.rs index 9f25647..f92c201 100644 --- a/examples/text-from-text-streaming.rs +++ b/examples/text-from-text-streaming.rs @@ -23,11 +23,11 @@ async fn main() -> Result<(), Box> { let request = GenerateContentRequest::builder().contents(prompt).build(); let mut queue = gemini - .generate_content_stream(&request, "gemini-2.0-flash-001") + .stream_generate_content(&request, "gemini-2.0-flash-001") .await?; - while let Some(Ok(response)) = queue.next().await { - println!("Response: {:?}", response); + while let Some(response) = queue.next().await { + println!("Response: {:?}", response?); } Ok(()) diff --git a/src/client.rs b/src/client.rs index 914612f..52c066b 100644 --- a/src/client.rs +++ b/src/client.rs @@ -1,309 +1,218 @@ -use crate::error::Result as GeminiResult; -use std::sync::Arc; -use std::vec; -use tokio_stream::{Stream, StreamExt}; - -use deadqueue::unlimited::Queue; -use reqwest_eventsource::{Event, EventSource}; -use tracing::error; - -use crate::dialogue::Message; -use crate::error::{Error, Result}; -use crate::prelude::{ - Candidate, Content, CountTokensRequest, CountTokensResponse, GenerateContentRequest, - GenerateContentResponse, GenerateContentResponseResult, TextEmbeddingRequest, - TextEmbeddingResponse, -}; -use crate::types::{PredictImageRequest, PredictImageResponse, Role}; -use crate::{prelude::Part, token_provider::TokenProvider}; - -pub static AUTH_SCOPE: &[&str] = &["https://www.googleapis.com/auth/cloud-platform"]; - -#[derive(Clone, Debug)] -pub struct GeminiClient { - token_provider: T, - client: reqwest::Client, - api_endpoint: String, - project_id: String, - location_id: String, -} - -unsafe impl Send for GeminiClient {} -unsafe impl Sync for GeminiClient {} - -impl GeminiClient { - pub fn new( - token_provider: T, - api_endpoint: String, - project_id: String, - location_id: String, - ) -> Self { - GeminiClient { - token_provider, - client: reqwest::Client::new(), - api_endpoint, - project_id, - location_id, - } - } - - pub async fn generate_content_stream( - &self, - request: &GenerateContentRequest, - model: &str, - ) -> Result>> { - let access_token = self.token_provider.get_token(AUTH_SCOPE).await.unwrap(); - let endpoint_url = format!( - "https://{}/v1beta1/projects/{}/locations/{}/publishers/google/models/{}:streamGenerateContent?alt=sse", self.api_endpoint, self.project_id, self.location_id, model, - ); - let client = self.client.clone(); - let request = request.clone(); - let req = client - .post(&endpoint_url) - .bearer_auth(access_token) - .json(&request); - - let event_source = EventSource::new(req).unwrap(); - - let mapped = event_source.filter_map(|event| { - let event = match event { - Ok(event) => event, - Err(reqwest_eventsource::Error::StreamEnded) => { - return Some(Err(Error::EventSourceClosedError)) - } - Err(e) => return Some(Err(e.into())), - }; - - let Event::Message(event_message) = event else { - return None; - }; - - let gemini_response: GenerateContentResponse = - match serde_json::from_str(&event_message.data) { - Ok(gemini_response) => gemini_response, - Err(e) => return Some(Err(e.into())), - }; - - let gemini_response = match gemini_response.into_result() { - Ok(gemini_response) => gemini_response, - Err(e) => return Some(Err(e)), - }; - - Some(Ok(gemini_response)) - }); - Ok(mapped) - } - - pub async fn stream_generate_content( - &self, - request: &GenerateContentRequest, - model: &str, - ) -> Arc>>> { - let queue = Arc::new(Queue::>>::new()); - let access_token = match self.token_provider.get_token(AUTH_SCOPE).await { - Ok(access_token) => access_token, - Err(e) => { - queue.push(Some(Err(e))); - return queue; - } - }; - - // Clone the queue and other necessary data to move into the async block. - let cloned_queue = queue.clone(); - let endpoint_url: String = format!( - "https://{}/v1beta1/projects/{}/locations/{}/publishers/google/models/{}:streamGenerateContent?alt=sse", self.api_endpoint, self.project_id, self.location_id, model, - ); - let client = self.client.clone(); - let request = request.clone(); - - // Start a thread to run the request in the background. - tokio::spawn(async move { - let req = client - .post(&endpoint_url) - .bearer_auth(access_token) - .json(&request); - - let mut event_source = match EventSource::new(req) { - Ok(event_source) => event_source, - Err(e) => { - cloned_queue.push(Some(Err(e.into()))); - return; - } - }; - while let Some(event) = event_source.next().await { - match event { - Ok(event) => { - if let Event::Message(event) = event { - let response: serde_json::error::Result = - serde_json::from_str(&event.data); - - match response { - Ok(response) => { - let result = response.into_result(); - let finished = match &result { - Ok(result) => result.candidates[0].finish_reason.is_some(), - Err(_) => true, - }; - cloned_queue.push(Some(result)); - if finished { - break; - } - } - Err(_) => { - tracing::error!("Error parsing message: {}", event.data); - break; - } - } - } - } - Err(e) => { - tracing::error!("Error in event source: {:?}", e); - break; - } - } - } - cloned_queue.push(None); - }); - - // Return the queue that will receive the responses. - queue - } - - pub async fn generate_content( - &self, - request: &GenerateContentRequest, - model: &str, - ) -> Result { - let access_token = self.token_provider.get_token(AUTH_SCOPE).await?; - let endpoint_url: String = format!( - "https://{}/v1beta1/projects/{}/locations/{}/publishers/google/models/{}:generateContent", self.api_endpoint, self.project_id, self.location_id, model, - ); - let resp = self - .client - .post(&endpoint_url) - .bearer_auth(access_token) - .json(&request) - .send() - .await?; - - let txt_json = resp.text().await?; - tracing::debug!("generate_content response: {:?}", txt_json); - match serde_json::from_str::(&txt_json) { - Ok(response) => Ok(response.into_result()?), - Err(e) => { - tracing::error!("Failed to parse response: {} with error {}", txt_json, e); - Err(e.into()) - } - } - } - - /// Prompts a conversation to the model. - pub async fn prompt_conversation(&self, messages: &[Message], model: &str) -> Result { - let request = GenerateContentRequest { - contents: messages - .iter() - .map(|m| Content { - role: Some(m.role), - parts: Some(vec![Part::Text(m.text.clone())]), - }) - .collect(), - generation_config: None, - tools: None, - system_instruction: None, - safety_settings: None, - }; - - let response = self.generate_content(&request, model).await?; - - // Check for errors in the response. - let mut candidates = GeminiClient::::collect_text_from_response(&response); - - match candidates.pop() { - Some(text) => Ok(Message::new(Role::Model, &text)), - None => Err(Error::NoCandidatesError), - } - } - - fn collect_text_from_response(response: &GenerateContentResponseResult) -> Vec { - response - .candidates - .iter() - .filter_map(Candidate::get_text) - .collect::>() - } - - pub async fn text_embeddings( - &self, - request: &TextEmbeddingRequest, - model: &str, - ) -> Result { - let endpoint_url = format!( - "https://{}/v1/projects/{}/locations/{}/publishers/google/models/{}:predict", - self.api_endpoint, self.project_id, self.location_id, model, - ); - let access_token = self.token_provider.get_token(AUTH_SCOPE).await?; - let resp = self - .client - .post(&endpoint_url) - .bearer_auth(access_token) - .json(&request) - .send() - .await?; - let txt_json = resp.text().await?; - tracing::debug!("text_embeddings response: {:?}", txt_json); - Ok(serde_json::from_str::(&txt_json)?) - } - - pub async fn count_tokens( - &self, - request: &CountTokensRequest, - model: &str, - ) -> Result { - let endpoint_url = format!( - "https://{}/v1beta1/projects/{}/locations/{}/publishers/google/models/{}:countTokens", - self.api_endpoint, self.project_id, self.location_id, model, - ); - let access_token = self.token_provider.get_token(AUTH_SCOPE).await?; - let resp = self - .client - .post(&endpoint_url) - .bearer_auth(access_token) - .json(&request) - .send() - .await?; - - let txt_json = resp.text().await?; - tracing::debug!("count_tokens response: {:?}", txt_json); - Ok(serde_json::from_str(&txt_json)?) - } - - pub async fn predict_image( - &self, - request: &PredictImageRequest, - model: &str, - ) -> Result { - let endpoint_url = format!( - "https://{}/v1/projects/{}/locations/{}/publishers/google/models/{}:predict", - self.api_endpoint, self.project_id, self.location_id, model, - ); - - let access_token = self.token_provider.get_token(AUTH_SCOPE).await?; - let resp = self - .client - .post(&endpoint_url) - .bearer_auth(access_token) - .json(&request) - .send() - .await?; - - let txt_json = resp.text().await?; - - match serde_json::from_str::(&txt_json) { - Ok(response) => Ok(response), - Err(e) => { - error!(response = txt_json, error = ?e, "Failed to parse response"); - Err(e.into()) - } - } - } -} +use std::vec; +use tracing::error; + +use tokio_stream::{Stream, StreamExt}; + +use crate::dialogue::Message; +use crate::error::{Error, Result}; +use crate::network::event_source::{EventSource, ServerSentEvent}; +use crate::prelude::{ + Candidate, Content, CountTokensRequest, CountTokensResponse, GenerateContentRequest, + GenerateContentResponse, GenerateContentResponseResult, TextEmbeddingRequest, + TextEmbeddingResponse, +}; +use crate::types::{PredictImageRequest, PredictImageResponse, Role}; +use crate::{prelude::Part, token_provider::TokenProvider}; + +pub const AUTH_SCOPE: &[&str] = &["https://www.googleapis.com/auth/cloud-platform"]; +const ENDPOINT_VERSION: &str = "v1beta1"; + +#[derive(Clone, Debug)] +pub struct GeminiClient { + token_provider: T, + client: reqwest::Client, + api_endpoint: String, + project_id: String, + location_id: String, +} + +impl GeminiClient { + pub fn new( + token_provider: T, + api_endpoint: String, + project_id: String, + location_id: String, + ) -> Self { + GeminiClient { + token_provider, + client: reqwest::Client::new(), + api_endpoint, + project_id, + location_id, + } + } + + pub async fn generate_content( + &self, + request: &GenerateContentRequest, + model: &str, + ) -> Result { + let access_token = self.token_provider.get_token(AUTH_SCOPE).await?; + let endpoint_url: String = format!( + "https://{}/{}/projects/{}/locations/{}/publishers/google/models/{}:generateContent", + self.api_endpoint, ENDPOINT_VERSION, self.project_id, self.location_id, model, + ); + let resp = self + .client + .post(&endpoint_url) + .bearer_auth(access_token) + .json(&request) + .send() + .await?; + + let txt_json = resp.text().await?; + tracing::debug!("generate_content response: {:?}", txt_json); + match serde_json::from_str::(&txt_json) { + Ok(response) => Ok(response.into_result()?), + Err(e) => { + tracing::error!("Failed to parse response: {} with error {}", txt_json, e); + Err(e.into()) + } + } + } + + /// Sends a content generation request and returns a stream of response chunks via SSE. + /// + /// Each item in the stream is a [`GenerateContentResponseResult`] containing one or more + /// candidates. Useful for displaying incremental output as it is generated. + pub async fn stream_generate_content( + &self, + request: &GenerateContentRequest, + model: &str, + ) -> Result> + use<'_, T>> { + let access_token = self.token_provider.get_token(AUTH_SCOPE).await?; + let endpoint_url = format!( + "https://{}/{}/projects/{}/locations/{}/publishers/google/models/{}:streamGenerateContent?alt=sse", + self.api_endpoint, ENDPOINT_VERSION, self.project_id, self.location_id, model, + ); + let response = self + .client + .post(&endpoint_url) + .bearer_auth(access_token) + .json(request) + .send() + .await?; + Ok(response.event_stream().filter_map(Self::parse_sse_event)) + } + + fn parse_sse_event( + event_result: std::result::Result, + ) -> Option> { + let data = event_result.map_err(Error::from).ok()?.data?; + Some( + serde_json::from_str::(&data) + .map_err(Error::from) + .and_then(|r| r.into_result()), + ) + } + + /// Prompts a conversation to the model. + pub async fn prompt_conversation(&self, messages: &[Message], model: &str) -> Result { + let request = GenerateContentRequest { + contents: messages + .iter() + .map(|m| Content { + role: Some(m.role), + parts: Some(vec![Part::text(m.text.clone())]), + }) + .collect(), + generation_config: None, + tools: None, + system_instruction: None, + safety_settings: None, + }; + + let response = self.generate_content(&request, model).await?; + + // Check for errors in the response. + let mut candidates = GeminiClient::::collect_text_from_response(&response); + + match candidates.pop() { + Some(text) => Ok(Message::new(Role::Model, &text)), + None => Err(Error::NoCandidatesError), + } + } + + fn collect_text_from_response(response: &GenerateContentResponseResult) -> Vec { + response + .candidates + .iter() + .filter_map(Candidate::get_text) + .collect::>() + } + + pub async fn text_embeddings( + &self, + request: &TextEmbeddingRequest, + model: &str, + ) -> Result { + let endpoint_url = format!( + "https://{}/v1/projects/{}/locations/{}/publishers/google/models/{}:predict", + self.api_endpoint, self.project_id, self.location_id, model, + ); + let access_token = self.token_provider.get_token(AUTH_SCOPE).await?; + let resp = self + .client + .post(&endpoint_url) + .bearer_auth(access_token) + .json(&request) + .send() + .await?; + let txt_json = resp.text().await?; + tracing::debug!("text_embeddings response: {:?}", txt_json); + Ok(serde_json::from_str::(&txt_json)?) + } + + pub async fn count_tokens( + &self, + request: &CountTokensRequest, + model: &str, + ) -> Result { + let endpoint_url = format!( + "https://{}/{}/projects/{}/locations/{}/publishers/google/models/{}:countTokens", + self.api_endpoint, ENDPOINT_VERSION, self.project_id, self.location_id, model, + ); + let access_token = self.token_provider.get_token(AUTH_SCOPE).await?; + let resp = self + .client + .post(&endpoint_url) + .bearer_auth(access_token) + .json(&request) + .send() + .await?; + + let txt_json = resp.text().await?; + tracing::debug!("count_tokens response: {:?}", txt_json); + Ok(serde_json::from_str(&txt_json)?) + } + + pub async fn predict_image( + &self, + request: &PredictImageRequest, + model: &str, + ) -> Result { + let endpoint_url = format!( + "https://{}/v1/projects/{}/locations/{}/publishers/google/models/{}:predict", + self.api_endpoint, self.project_id, self.location_id, model, + ); + + let access_token = self.token_provider.get_token(AUTH_SCOPE).await?; + let resp = self + .client + .post(&endpoint_url) + .bearer_auth(access_token) + .json(&request) + .send() + .await?; + + let txt_json = resp.text().await?; + + match serde_json::from_str::(&txt_json) { + Ok(response) => Ok(response), + Err(e) => { + error!(response = txt_json, error = ?e, "Failed to parse response"); + Err(e.into()) + } + } + } +} diff --git a/src/dialogue.rs b/src/dialogue.rs index e819494..7aab87f 100644 --- a/src/dialogue.rs +++ b/src/dialogue.rs @@ -1,46 +1,46 @@ -use serde::{Deserialize, Serialize}; - -use crate::{client::GeminiClient, error::Result, prelude::TokenProvider, types::Role}; - -#[derive(Clone, Debug, Serialize, Deserialize)] -pub struct Message { - pub role: Role, - pub text: String, -} - -impl Message { - pub fn new(role: Role, text: &str) -> Self { - Message { - role, - text: text.to_string(), - } - } -} - -#[derive(Clone, Debug, Serialize, Deserialize)] -pub struct Dialogue { - model: String, - messages: Vec, -} - -impl Dialogue { - pub fn new(model: &str) -> Self { - Dialogue { - model: model.to_string(), - messages: vec![], - } - } - - pub async fn do_turn( - &mut self, - gemini: &GeminiClient, - message: &str, - ) -> Result { - self.messages.push(Message::new(Role::User, message)); - let response = gemini - .prompt_conversation(&self.messages, &self.model) - .await?; - self.messages.push(response.clone()); - Ok(response) - } -} +use serde::{Deserialize, Serialize}; + +use crate::{client::GeminiClient, error::Result, prelude::TokenProvider, types::Role}; + +#[derive(Clone, Debug, Serialize, Deserialize)] +pub struct Message { + pub role: Role, + pub text: String, +} + +impl Message { + pub fn new(role: Role, text: &str) -> Self { + Message { + role, + text: text.to_string(), + } + } +} + +#[derive(Clone, Debug, Serialize, Deserialize)] +pub struct Dialogue { + model: String, + messages: Vec, +} + +impl Dialogue { + pub fn new(model: &str) -> Self { + Dialogue { + model: model.to_string(), + messages: vec![], + } + } + + pub async fn do_turn( + &mut self, + gemini: &GeminiClient, + message: &str, + ) -> Result { + self.messages.push(Message::new(Role::User, message)); + let response = gemini + .prompt_conversation(&self.messages, &self.model) + .await?; + self.messages.push(response.clone()); + Ok(response) + } +} diff --git a/src/error.rs b/src/error.rs index f006510..a5953fa 100644 --- a/src/error.rs +++ b/src/error.rs @@ -1,90 +1,71 @@ -use std::fmt::Display; - -use reqwest_eventsource::CannotCloneRequestError; - -use crate::types; - -pub type Result = std::result::Result; - -#[derive(Debug)] -pub enum Error { - Env(std::env::VarError), - HttpClient(reqwest::Error), - Token(gcp_auth::Error), - Serde(serde_json::Error), - VertexError(types::VertexApiError), - NoCandidatesError, - CannotCloneRequestError(CannotCloneRequestError), - EventSourceError(reqwest_eventsource::Error), - EventSourceClosedError, -} - -impl Display for Error { - fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { - match &self { - Error::Env(e) => write!(f, "Environment variable error: {}", e), - Error::HttpClient(e) => write!(f, "HTTP Client error: {}", e), - Error::Token(e) => write!(f, "Token error: {}", e), - Error::Serde(e) => write!(f, "Serde error: {}", e), - Error::VertexError(e) => { - write!(f, "Vertex error: {}", e) - } - Error::NoCandidatesError => { - write!(f, "No candidates returned for the prompt") - } - Error::CannotCloneRequestError(e) => { - write!(f, "Cannot clone request: {}", e) - } - Error::EventSourceError(e) => { - write!(f, "EventSourrce Error: {}", e) - } - Error::EventSourceClosedError => { - write!(f, "EventSource closed error") - } - } - } -} - -impl std::error::Error for Error {} - -impl From for Error { - fn from(e: reqwest::Error) -> Self { - Error::HttpClient(e) - } -} - -impl From for Error { - fn from(e: std::env::VarError) -> Self { - Error::Env(e) - } -} - -impl From for Error { - fn from(e: gcp_auth::Error) -> Self { - Error::Token(e) - } -} - -impl From for Error { - fn from(e: serde_json::Error) -> Self { - Error::Serde(e) - } -} - -impl From for Error { - fn from(e: types::VertexApiError) -> Self { - Error::VertexError(e) - } -} - -impl From for Error { - fn from(e: CannotCloneRequestError) -> Self { - Error::CannotCloneRequestError(e) - } -} - -impl From for Error { - fn from(e: reqwest_eventsource::Error) -> Self { - Error::EventSourceError(e) - } -} +use std::fmt::Display; + +use tokio_util::codec::LinesCodecError; + +use crate::types; + +pub type Result = std::result::Result; + +#[derive(Debug)] +pub enum Error { + Env(std::env::VarError), + HttpClient(reqwest::Error), + Token(gcp_auth::Error), + Serde(serde_json::Error), + VertexError(types::VertexApiError), + NoCandidatesError, + /// An error occurred while decoding the SSE event stream. + EventSourceError(LinesCodecError), +} + +impl Display for Error { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + match &self { + Error::Env(e) => write!(f, "Environment variable error: {e}"), + Error::HttpClient(e) => write!(f, "HTTP Client error: {e}"), + Error::Token(e) => write!(f, "Token error: {e}"), + Error::Serde(e) => write!(f, "Serde error: {e}"), + Error::VertexError(e) => write!(f, "Vertex error: {e}"), + Error::NoCandidatesError => write!(f, "No candidates returned for the prompt"), + Error::EventSourceError(e) => write!(f, "EventSource error: {e}"), + } + } +} + +impl std::error::Error for Error {} + +impl From for Error { + fn from(e: reqwest::Error) -> Self { + Error::HttpClient(e) + } +} + +impl From for Error { + fn from(e: std::env::VarError) -> Self { + Error::Env(e) + } +} + +impl From for Error { + fn from(e: serde_json::Error) -> Self { + Error::Serde(e) + } +} + +impl From for Error { + fn from(e: gcp_auth::Error) -> Self { + Error::Token(e) + } +} + +impl From for Error { + fn from(e: types::VertexApiError) -> Self { + Error::VertexError(e) + } +} + +impl From for Error { + fn from(e: LinesCodecError) -> Self { + Error::EventSourceError(e) + } +} diff --git a/src/lib.rs b/src/lib.rs index 25b517c..ce61cf3 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -1,6 +1,7 @@ mod client; mod dialogue; pub mod error; +pub(crate) mod network; mod token_provider; mod types; diff --git a/src/network/event_source.rs b/src/network/event_source.rs new file mode 100644 index 0000000..615dd58 --- /dev/null +++ b/src/network/event_source.rs @@ -0,0 +1,108 @@ +use reqwest::Response; +use std::mem; +use tokio_stream::{Stream, StreamExt}; +use tokio_util::{ + codec::{Decoder, FramedRead, LinesCodec, LinesCodecError}, + io::StreamReader, +}; +use tracing::warn; + +static EVENT: &str = "event: "; +static DATA: &str = "data: "; +static ID: &str = "id: "; +static RETRY: &str = "retry: "; + +/// Extension trait for converting an HTTP response into a stream of [`ServerSentEvent`]s. +pub trait EventSource { + /// Consumes the response and returns a stream of parsed SSE events. + fn event_stream(self) -> impl Stream>; +} + +impl EventSource for Response { + fn event_stream(self) -> impl Stream> { + stream_response(self) + } +} + +/// A parsed Server-Sent Event. +#[derive(Debug, Default, Clone)] +pub struct ServerSentEvent { + pub event: Option, + /// The event payload (from one or more `data:` fields, joined by `\n`). + pub data: Option, + pub id: Option, + pub retry: Option, +} + +/// [`Decoder`] that parses a byte stream of SSE-formatted data into [`ServerSentEvent`]s. +pub struct ServerSentEventsCodec { + lines_codec: LinesCodec, + next: ServerSentEvent, +} + +impl Default for ServerSentEventsCodec { + fn default() -> Self { + Self::new() + } +} + +impl ServerSentEventsCodec { + pub fn new() -> Self { + Self { + lines_codec: LinesCodec::new(), + next: Default::default(), + } + } +} + +impl Decoder for ServerSentEventsCodec { + type Item = ServerSentEvent; + type Error = LinesCodecError; + + fn decode( + &mut self, + src: &mut tokio_util::bytes::BytesMut, + ) -> Result, Self::Error> { + let Some(mut line) = self.lines_codec.decode(src)? else { + return Ok(None); + }; + + if line.is_empty() { + return Ok(Some(mem::take(&mut self.next))); + } + + if line.starts_with(EVENT) { + line.drain(..EVENT.len()); + self.next.event = Some(line); + } else if line.starts_with(DATA) { + line.drain(..DATA.len()); + if let Some(ref mut existing) = self.next.data { + existing.push('\n'); + existing.push_str(&line); + } else { + self.next.data = Some(line); + } + } else if line.starts_with(ID) { + line.drain(..ID.len()); + self.next.id = Some(line); + } else if line.starts_with(RETRY) { + line.drain(..RETRY.len()); + let Ok(retry) = line.parse() else { + warn!(line, "Received invalid retry value"); + return Ok(None); + }; + self.next.retry = Some(retry); + } + + Ok(None) + } +} + +/// Converts a [`Response`] into a stream of [`ServerSentEvent`]s. +pub fn stream_response( + response: Response, +) -> impl Stream> { + let bytes_stream = response.bytes_stream(); + let body_reader = StreamReader::new(bytes_stream.map(|r| r.map_err(std::io::Error::other))); + FramedRead::new(body_reader, ServerSentEventsCodec::new()) +} diff --git a/src/network/mod.rs b/src/network/mod.rs new file mode 100644 index 0000000..aafb61f --- /dev/null +++ b/src/network/mod.rs @@ -0,0 +1 @@ +pub mod event_source; diff --git a/src/token_provider.rs b/src/token_provider.rs index 880ae79..7583193 100644 --- a/src/token_provider.rs +++ b/src/token_provider.rs @@ -1,18 +1,18 @@ -use std::sync::Arc; - -use crate::error::Result; - -pub trait TokenProvider { - fn get_token(&self, scope: &[&str]) - -> impl std::future::Future> + Send; -} - -impl TokenProvider for Arc { - async fn get_token(&self, scope: &[&str]) -> Result { - let token = self.token(scope).await; - match token { - Ok(token) => Ok(token.as_str().to_string()), - Err(e) => Err(e.into()), - } - } -} +use std::sync::Arc; + +use crate::error::Result; + +pub trait TokenProvider { + fn get_token(&self, scope: &[&str]) + -> impl std::future::Future> + Send; +} + +impl TokenProvider for Arc { + async fn get_token(&self, scope: &[&str]) -> Result { + let token = self.token(scope).await; + match token { + Ok(token) => Ok(token.as_str().to_string()), + Err(e) => Err(e.into()), + } + } +} diff --git a/src/types/common.rs b/src/types/common.rs index e5ed019..3c08ff5 100644 --- a/src/types/common.rs +++ b/src/types/common.rs @@ -1,6 +1,7 @@ -use std::{collections::HashMap, fmt::Display, str::FromStr, vec}; +use std::{fmt::Display, str::FromStr, vec}; use serde::{Deserialize, Serialize}; +use serde_json::Value; #[derive(Clone, Default, Debug, Serialize, Deserialize)] pub struct Content { @@ -13,10 +14,7 @@ impl Content { self.parts.as_ref().map(|parts| { parts .iter() - .filter_map(|part| match part { - Part::Text(text) => Some(text.clone()), - _ => None, - }) + .filter_map(|part| part.text.clone()) .collect::() }) } @@ -26,6 +24,27 @@ impl Content { } } +pub struct SystemPrompt { + content: Content, +} + +impl> From for SystemPrompt { + fn from(value: T) -> Self { + SystemPrompt { + content: Content { + parts: Some(vec![Part::text(value.as_ref())]), + role: None, + }, + } + } +} + +impl From for Content { + fn from(value: SystemPrompt) -> Self { + value.content + } +} + #[derive(Default)] pub struct ContentBuilder { content: Content, @@ -33,7 +52,7 @@ pub struct ContentBuilder { impl ContentBuilder { pub fn add_text_part>(self, text: T) -> Self { - self.add_part(Part::Text(text.into())) + self.add_part(Part::text(text)) } pub fn add_part(mut self, part: Part) -> Self { @@ -54,7 +73,7 @@ impl ContentBuilder { } } -#[derive(Clone, Copy, Debug, Serialize, Deserialize)] +#[derive(Clone, Copy, Debug, PartialEq, Eq, Serialize, Deserialize)] #[serde(rename_all = "lowercase")] pub enum Role { User, @@ -83,83 +102,83 @@ impl FromStr for Role { } } -#[derive(Clone, Debug, Serialize)] +#[derive(Clone, Debug, Default, Serialize, Deserialize)] #[serde(rename_all = "camelCase")] -pub enum Part { - Text(String), - InlineData { - mime_type: String, - data: String, - }, - FileData { - mime_type: String, - file_uri: String, - }, - FunctionCall { - name: String, - args: HashMap, - }, +pub struct Part { + #[serde(skip_serializing_if = "Option::is_none")] + pub text: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub inline_data: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub file_data: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub function_call: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub function_response: Option, + /// Opaque signature emitted by Gemini 2.5+ thinking models on function-call + /// parts. Must be echoed back on the same part — and on the corresponding + /// `functionResponse` part — in subsequent turns, otherwise Vertex rejects + /// the request with "function call X is missing a thought_signature". + #[serde(skip_serializing_if = "Option::is_none")] + pub thought_signature: Option, + /// `Some(true)` on a response part marks it as a reasoning/thinking token + /// (only emitted when `thinkingConfig.includeThoughts` is set on the + /// request). Absent on regular text and on outgoing requests. + #[serde(skip_serializing_if = "Option::is_none")] + pub thought: Option, } -#[derive(Deserialize)] -#[serde(rename_all = "camelCase")] -struct PartHelper { - text: Option, - inline_data: Option, - file_data: Option, - function_call: Option, -} - -#[derive(Deserialize)] -#[serde(rename_all = "camelCase")] -struct InlineDataHelper { - mime_type: String, - data: String, -} - -#[derive(Deserialize)] -#[serde(rename_all = "camelCase")] -struct FileDataHelper { - mime_type: String, - file_uri: String, -} - -#[derive(Deserialize)] -#[serde(rename_all = "camelCase")] -struct FunctionCallHelper { - name: String, - args: HashMap, -} - -impl<'de> Deserialize<'de> for Part { - fn deserialize(deserializer: D) -> Result - where - D: serde::Deserializer<'de>, - { - let helper = PartHelper::deserialize(deserializer)?; - if let Some(text) = helper.text { - return Ok(Part::Text(text)); +impl Part { + pub fn text(text: impl Into) -> Self { + Part { + text: Some(text.into()), + ..Default::default() } - if let Some(inline_data) = helper.inline_data { - return Ok(Part::InlineData { - mime_type: inline_data.mime_type, - data: inline_data.data, - }); + } + + pub fn function_call(name: impl Into, args: Value) -> Self { + Part { + function_call: Some(FunctionCallData { + name: name.into(), + args, + }), + ..Default::default() } - if let Some(file_data) = helper.file_data { - return Ok(Part::FileData { - mime_type: file_data.mime_type, - file_uri: file_data.file_uri, - }); + } + + pub fn function_response(name: impl Into, response: Value) -> Self { + Part { + function_response: Some(FunctionResponseData { + name: name.into(), + response, + }), + ..Default::default() } - if let Some(function_call) = helper.function_call { - return Ok(Part::FunctionCall { - name: function_call.name, - args: function_call.args, - }); - } - Err(serde::de::Error::custom( - "Part does not contain any recognizable variant", - )) } } + +#[derive(Clone, Debug, Serialize, Deserialize)] +pub struct FunctionCallData { + pub name: String, + pub args: Value, +} + +#[derive(Clone, Debug, Serialize, Deserialize)] +pub struct FunctionResponseData { + pub name: String, + pub response: Value, +} + +#[derive(Clone, Debug, Serialize, Deserialize)] +#[serde(rename_all = "camelCase")] +pub struct InlineData { + pub mime_type: String, + pub data: String, +} + +#[derive(Clone, Debug, Serialize, Deserialize)] +#[serde(rename_all = "camelCase")] +pub struct FileData { + pub mime_type: String, + pub file_uri: String, +} diff --git a/src/types/count_tokens.rs b/src/types/count_tokens.rs index b586339..3284574 100644 --- a/src/types/count_tokens.rs +++ b/src/types/count_tokens.rs @@ -22,7 +22,7 @@ impl CountTokensRequestBuilder { pub fn from_prompt(prompt: &str) -> Self { CountTokensRequestBuilder { contents: Content { - parts: Some(vec![super::Part::Text(prompt.to_string())]), + parts: Some(vec![super::Part::text(prompt)]), ..Default::default() }, } diff --git a/src/types/generate_content.rs b/src/types/generate_content.rs index ada928d..37fe674 100644 --- a/src/types/generate_content.rs +++ b/src/types/generate_content.rs @@ -1,5 +1,3 @@ -use std::collections::HashMap; - use serde::{Deserialize, Serialize}; use serde_json::Value; @@ -125,6 +123,37 @@ pub struct GenerationConfig { pub response_mime_type: Option, #[serde(skip_serializing_if = "Option::is_none")] pub response_schema: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub thinking_config: Option, +} + +/// Configures Gemini 2.5+ thinking models. See +/// . +#[derive(Clone, Debug, Serialize, Deserialize, Default)] +#[serde(rename_all = "camelCase")] +pub struct ThinkingConfig { + /// When `true`, the response includes parts with `thought: true` carrying + /// the model's reasoning tokens. + #[serde(skip_serializing_if = "Option::is_none")] + pub include_thoughts: Option, + /// Maximum tokens the model may spend thinking. `Some(0)` disables + /// thinking entirely on models that support it. + #[serde(skip_serializing_if = "Option::is_none")] + pub thinking_budget: Option, + /// Coarse-grained thinking effort. Mutually-exclusive shorthand for + /// `thinking_budget` on supported models. + #[serde(skip_serializing_if = "Option::is_none")] + pub thinking_level: Option, +} + +/// Coarse-grained thinking effort. Serialized as `SCREAMING_SNAKE_CASE` to +/// match the Vertex API enum. +#[derive(Clone, Copy, Debug, Serialize, Deserialize)] +#[serde(rename_all = "SCREAMING_SNAKE_CASE")] +pub enum ThinkingLevel { + ThinkingLevelUnspecified, + Low, + High, } impl GenerationConfig { @@ -184,6 +213,11 @@ impl GenerationConfigBuilder { self } + pub fn thinking_config(mut self, thinking_config: ThinkingConfig) -> Self { + self.generation_config.thinking_config = Some(thinking_config); + self + } + pub fn build(self) -> GenerationConfig { self.generation_config } @@ -294,22 +328,7 @@ pub struct UsageMetadata { pub struct FunctionDeclaration { pub name: String, pub description: String, - pub parameters: FunctionParameters, -} - -#[derive(Clone, Debug, Serialize, Deserialize)] -#[serde(rename_all = "camelCase")] -pub struct FunctionParameters { - pub r#type: String, - pub properties: HashMap, - pub required: Vec, -} - -#[derive(Clone, Debug, Serialize, Deserialize)] -#[serde(rename_all = "camelCase")] -pub struct FunctionParametersProperty { - pub r#type: String, - pub description: String, + pub parameters: Value, } #[derive(Debug, Serialize, Deserialize)] @@ -485,6 +504,47 @@ mod tests { serde_json::from_str::(input).unwrap(); } + #[test] + fn parses_function_call_with_thought_signature() { + let input = r#"{ + "candidates": [ + { + "content": { + "role": "model", + "parts": [ + { + "functionCall": { + "name": "delegate_to_agent", + "args": { "agent_name": "general", "task": "Say hello." } + }, + "thoughtSignature": "CsYCAY89a18IQwmp" + } + ] + }, + "finishReason": "STOP" + } + ], + "usageMetadata": { "promptTokenCount": 112, "totalTokenCount": 196 } + }"#; + let resp = serde_json::from_str::(input).unwrap(); + if let GenerateContentResponse::Ok(r) = resp { + let part = &r.candidates[0] + .content + .as_ref() + .unwrap() + .parts + .as_ref() + .unwrap()[0]; + assert_eq!( + part.function_call.as_ref().unwrap().name, + "delegate_to_agent" + ); + assert_eq!(part.thought_signature.as_deref(), Some("CsYCAY89a18IQwmp")); + } else { + panic!("expected Ok response"); + } + } + #[test] fn parses_safety_rating_without_scores() { let input = r#"{ @@ -531,34 +591,4 @@ mod tests { }"#; serde_json::from_str::(input).unwrap(); } - - #[test] - fn parses_response_with_thought_signature() { - let input = r#"{ - "candidates": [ - { - "content": { - "role": "model", - "parts": [ - { - "text": "Hello world", - "thoughtSignature": "AY89a1+twFsQ+oYDReN5PU76yp0ciDinNh3MgyCbe/BLvyh93lje7qHOCJuNxaPULnqmgvKvtuQnjyf2wwD20Dl2rWnxbdZZAeqhRJEdvKOc4LQ=" - } - ] - }, - "finishReason": "STOP" - } - ], - "usageMetadata": { - "promptTokenCount": 877, - "candidatesTokenCount": 547, - "totalTokenCount": 1424, - "trafficType": "ON_DEMAND" - }, - "modelVersion": "gemini-3.5-flash-lite", - "createTime": "2026-09-06T10:36:16.330983Z", - "responseId": "IEKdaueZFKyrz_IPp8Sc4QQ" - }"#; - serde_json::from_str::(input).unwrap(); - } } diff --git a/src/types/predict_image.rs b/src/types/predict_image.rs index 3531a38..a33cbaf 100644 --- a/src/types/predict_image.rs +++ b/src/types/predict_image.rs @@ -12,6 +12,10 @@ pub struct PredictImageRequest { pub struct PredictImageRequestPrompt { /// The text prompt for the image. /// The following models support different values for this parameter: + /// - `imagen-4.0-generate-001`: up to 480 tokens. + /// - `imagen-4.0-fast-generate-001`: up to 480 tokens. + /// - `imagen-4.0-ultra-generate-001`: up to 480 tokens. + /// - `imagen-3.0-generate-002`: up to 480 tokens. /// - `imagen-3.0-generate-001`: up to 480 tokens. /// - `imagen-3.0-fast-generate-001`: up to 480 tokens. /// - `imagegeneration@006`: up to 128 tokens. @@ -24,27 +28,20 @@ pub struct PredictImageRequestPrompt { #[serde(rename_all = "camelCase")] pub struct PredictImageRequestParameters { /// The number of images to generate. The default value is 4. - /// The following models support different values for this parameter: - /// - `imagen-3.0-generate-001`: 1 to 4. - /// - `imagen-3.0-fast-generate-001`: 1 to 4. - /// - `imagegeneration@006`: 1 to 4. - /// - `imagegeneration@005`: 1 to 4. - /// - `imagegeneration@002`: 1 to 8. + /// Supported range: 1 to 4 for all current models. pub sample_count: i32, - /// The random seed for image generation. This is not available when addWatermark is set to - /// true. + /// The random seed for image generation. This is not available when `add_watermark` is set to + /// `true`. #[serde(skip_serializing_if = "Option::is_none")] pub seed: Option, - /// Optional. An optional parameter to use an LLM-based prompt rewriting feature to deliver - /// higher quality images that better reflect the original prompt's intent. Disabling this - /// feature may impact image quality and prompt adherence - #[serde(skip_serializing_if = "Option::is_none")] - pub enhance_prompt: Option, - /// A description of what to discourage in the generated images. /// The following models support this parameter: + /// - `imagen-4.0-generate-001`: up to 480 tokens. + /// - `imagen-4.0-fast-generate-001`: up to 480 tokens. + /// - `imagen-4.0-ultra-generate-001`: up to 480 tokens. + /// - `imagen-3.0-generate-002`: up to 480 tokens. /// - `imagen-3.0-generate-001`: up to 480 tokens. /// - `imagen-3.0-fast-generate-001`: up to 480 tokens. /// - `imagegeneration@006`: up to 128 tokens. @@ -53,13 +50,8 @@ pub struct PredictImageRequestParameters { #[serde(skip_serializing_if = "Option::is_none")] pub negative_prompt: Option, - /// The aspect ratio for the image. The default value is "1:1". - /// The following models support different values for this parameter: - /// - `imagen-3.0-generate-001`: "1:1", "9:16", "16:9", "3:4", or "4:3". - /// - `imagen-3.0-fast-generate-001`: "1:1", "9:16", "16:9", "3:4", or "4:3". - /// - `imagegeneration@006`: "1:1", "9:16", "16:9", "3:4", or "4:3". - /// - `imagegeneration@005`: "1:1" or "9:16". - /// - `imagegeneration@002`: "1:1". + /// The aspect ratio for the image. The default value is `"1:1"`. + /// Supported values: `"1:1"`, `"9:16"`, `"16:9"`, `"3:4"`, `"4:3"`. #[serde(skip_serializing_if = "Option::is_none")] pub aspect_ratio: Option, @@ -77,7 +69,7 @@ pub struct PredictImageRequestParameters { /// - "cyberpunk" /// - "pop_art" /// - /// Pre-defined styles is only supported for model imagegeneration@002 + /// **Deprecated**: only supported by the legacy `imagegeneration@002` model. #[serde(skip_serializing_if = "Option::is_none")] pub sample_image_style: Option, @@ -87,56 +79,41 @@ pub struct PredictImageRequestParameters { /// - `"allow_all"`: Allow generation of people of all ages. /// /// The default value is `"allow_adult"`. - /// - /// Supported by the models `imagen-3.0-generate-001`, `imagen-3.0-fast-generate-001`, and - /// `imagegeneration@006` only. #[serde(skip_serializing_if = "Option::is_none")] - pub person_generation: Option, - - /// Optional. The language code that corresponds to your text prompt language. - /// The following values are supported: - /// - auto: Automatic detection. If Imagen detects a supported language, the prompt and an - /// optional negative prompt are translated to English. If the language detected isn't - /// supported, Imagen uses the input text verbatim, which might result in an unexpected - /// output. No error code is returned. - /// - en: English (if omitted, the default value) - /// - zh or zh-CN: Chinese (simplified) - /// - zh-TW: Chinese (traditional) - /// - hi: Hindi - /// - ja: Japanese - /// - ko: Korean - /// - pt: Portuguese - /// - es: Spanish - #[serde(skip_serializing_if = "Option::is_none")] - pub language: Option, + pub person_generation: Option, /// Adds a filter level to safety filtering. The following values are supported: - /// - /// - "block_low_and_above": Strongest filtering level, most strict blocking. - /// Deprecated value: "block_most". - /// - "block_medium_and_above": Block some problematic prompts and responses. - /// Deprecated value: "block_some". - /// - "block_only_high": Reduces the number of requests blocked due to safety filters. May - /// increase objectionable content generated by Imagen. Deprecated value: "block_few". - /// - "block_none": Block very few problematic prompts and responses. Access to this feature - /// is restricted. Previous field value: "block_fewest". - /// - /// The default value is "block_medium_and_above". - /// - /// Supported by the models `imagen-3.0-generate-001`, `imagen-3.0-fast-generate-001`, and - /// `imagegeneration@006` only. + /// - `"block_low_and_above"`: Strongest filtering, most strict blocking. + /// - `"block_medium_and_above"`: Block medium and high severity content. Default. + /// - `"block_only_high"`: Only block high severity content. + /// - `"block_none"`: No blocking. #[serde(skip_serializing_if = "Option::is_none")] - pub safety_setting: Option, + pub safety_setting: Option, - /// Add an invisible watermark to the generated images. The default value is `false` for the - /// `imagegeneration@002` and `imagegeneration@005` models, and `true` for the - /// `imagen-3.0-fast-generate-001`, `imagegeneration@006`, and imagegeneration@006 models. + /// Add an invisible SynthID watermark to the generated images. Defaults to `true` for + /// Imagen 3 and 4 models. Not available when `seed` is set. #[serde(skip_serializing_if = "Option::is_none")] pub add_watermark: Option, /// Cloud Storage URI to store the generated images. #[serde(skip_serializing_if = "Option::is_none")] pub storage_uri: Option, + + /// Enables LLM-based prompt rewriting to improve image quality. Default: `true`. + /// Recommended to set to `false` for `imagen-4.0-fast-generate-001` when the prompt is + /// complex. + #[serde(skip_serializing_if = "Option::is_none")] + pub enhance_prompt: Option, + + /// Language of the prompt. Default: `"auto"`. + /// Supported values: `"auto"`, `"en"`, `"zh"`, `"zh-CN"`, `"zh-TW"`, `"hi"`, `"ja"`, + /// `"ko"`, `"pt"`, `"es"`. + #[serde(skip_serializing_if = "Option::is_none")] + pub language: Option, + + /// Output image resolution. `"1K"` (default) or `"2K"`. + #[serde(skip_serializing_if = "Option::is_none")] + pub sample_image_size: Option, } #[derive(Debug, Serialize, Deserialize)] @@ -144,10 +121,9 @@ pub struct PredictImageRequestParameters { pub struct PredictImageRequestParametersOutputOptions { /// The image format that the output should be saved as. The following values are supported: /// - /// - "image/png": Save as a PNG image - /// - "image/jpeg": Save as a JPEG image - /// - /// The default value is "image/png".v + /// - `"image/png"`: Save as a PNG image (default) + /// - `"image/jpeg"`: Save as a JPEG image + /// - `"image/webp"`: Save as a WebP image pub mime_type: Option, /// The level of compression if the output type is "image/jpeg". @@ -168,20 +144,3 @@ pub struct PredictImageResponsePrediction { pub bytes_base64_encoded: Vec, pub mime_type: String, } - -#[derive(Debug, Serialize, Deserialize)] -#[serde(rename_all = "snake_case")] -pub enum PersonGeneration { - DontAllow, - AllowAdult, - AllowAll, -} - -#[derive(Debug, Serialize, Deserialize)] -#[serde(rename_all = "snake_case")] -pub enum PredictImageSafetySetting { - BlockLowAndAbove, - BlockMediumAndAbove, - BlockOnlyHigh, - BlockNone, -}