feat: add SSE streaming, thinking models, and modernized Part struct
Some checks failed
Rust / build (push) Has been cancelled

This commit is contained in:
2026-09-20 16:42:55 +01:00
parent ad82fc712b
commit 494aab6d14
18 changed files with 703 additions and 696 deletions

View File

@@ -1,26 +1,26 @@
[package] [package]
name = "gemini-rs" name = "gemini-rs"
version = "0.1.0" version = "0.1.0"
edition = "2021" edition = "2024"
# See more keys and their definitions at https://doc.rust-lang.org/cargo/reference/manifest.html # See more keys and their definitions at https://doc.rust-lang.org/cargo/reference/manifest.html
[dependencies] [dependencies]
deadqueue = "0.2" deadqueue = "0.2"
gcp_auth = "0.12" gcp_auth = "0.12"
reqwest = { version = "0.12", features = ["json", "gzip"] } reqwest = { version = "0.13", features = ["json", "gzip", "stream"] }
reqwest-eventsource = "0.6" tokio-util = "0.7"
serde = { version = "1", features = ["derive"] } serde = { version = "1", features = ["derive"] }
serde_json = { version = "1"} serde_json = { version = "1" }
serde_with = { version = "3.9", features = ["base64"]} serde_with = { version = "3.21", features = ["base64"] }
tracing = "0.1" tracing = "0.1"
tokio = { version = "1" } tokio = { version = "1" }
tokio-stream = "0.1.17" tokio-stream = "0.1.18"
[dev-dependencies] [dev-dependencies]
console = "0.15.8" console = "0.16.4"
dialoguer = "0.11.0" dialoguer = "0.12.0"
image = "0.25.2" image = "0.25.10"
indicatif = "0.17.8" indicatif = "0.18.6"
tokio = { version = "1.37.0", features = ["full"] } tokio = { version = "1.52.3", features = ["full"] }
tracing-subscriber = "0.3.18" tracing-subscriber = "0.3.23"

View File

@@ -1,9 +1,8 @@
use std::{error::Error, io::Cursor}; use std::{error::Error, io::Cursor};
use gemini_rs::prelude::{ use gemini_rs::prelude::{
GeminiClient, PersonGeneration, PredictImageRequest, PredictImageRequestParameters, GeminiClient, PredictImageRequest, PredictImageRequestParameters,
PredictImageRequestParametersOutputOptions, PredictImageRequestPrompt, PredictImageRequestParametersOutputOptions, PredictImageRequestPrompt,
PredictImageSafetySetting,
}; };
use image::{ImageFormat, ImageReader}; use image::{ImageFormat, ImageReader};
@@ -35,8 +34,8 @@ pub async fn main() -> Result<(), Box<dyn Error>> {
mime_type: Some("image/jpeg".to_string()), mime_type: Some("image/jpeg".to_string()),
compression_quality: Some(75), compression_quality: Some(75),
}), }),
person_generation: Some(PersonGeneration::AllowAll), person_generation: Some("allow_all".to_string()),
safety_setting: Some(PredictImageSafetySetting::BlockOnlyHigh), safety_setting: Some("block_only_high".to_string()),
..Default::default() ..Default::default()
}, },
}; };

View File

@@ -20,7 +20,7 @@ async fn main() -> Result<(), Box<dyn std::error::Error>> {
let request = GenerateContentRequest { let request = GenerateContentRequest {
contents: vec![Content { contents: vec![Content {
role: Some(Role::User), role: Some(Role::User),
parts: Some(vec![Part::Text(prompt.to_string())]), parts: Some(vec![Part::text(prompt)]),
}], }],
tools: Some(vec![Tools { tools: Some(vec![Tools {
google_search_retrieval: Some(GoogleSearchRetrieval::default()), google_search_retrieval: Some(GoogleSearchRetrieval::default()),

View File

@@ -20,7 +20,7 @@ async fn main() -> Result<(), Box<dyn std::error::Error>> {
let request = GenerateContentRequest { let request = GenerateContentRequest {
contents: vec![Content { contents: vec![Content {
role: Some(Role::User), role: Some(Role::User),
parts: Some(vec![Part::Text(prompt.to_string())]), parts: Some(vec![Part::text(prompt)]),
}], }],
tools: Some(vec![Tools { tools: Some(vec![Tools {
google_search: Some(GoogleSearch::default()), google_search: Some(GoogleSearch::default()),

View File

@@ -19,7 +19,7 @@ async fn main() -> Result<(), Box<dyn std::error::Error>> {
let request = GenerateContentRequest { let request = GenerateContentRequest {
contents: vec![Content { contents: vec![Content {
role: Some(Role::User), role: Some(Role::User),
parts: Some(vec![Part::Text(prompt.to_string())]), parts: Some(vec![Part::text(prompt)]),
}], }],
generation_config: Some(GenerationConfig { generation_config: Some(GenerationConfig {
response_mime_type: Some("application/json".to_string()), response_mime_type: Some("application/json".to_string()),

View File

@@ -21,7 +21,7 @@ async fn main() -> Result<(), Box<dyn std::error::Error>> {
let request = GenerateContentRequest { let request = GenerateContentRequest {
contents: vec![Content { contents: vec![Content {
role: Some(Role::User), role: Some(Role::User),
parts: Some(vec![Part::Text(prompt.to_string())]), parts: Some(vec![Part::text(prompt)]),
}], }],
safety_settings: Some(vec![SafetySetting { safety_settings: Some(vec![SafetySetting {
category: HarmCategory::HateSpeech, category: HarmCategory::HateSpeech,

View File

@@ -23,11 +23,11 @@ async fn main() -> Result<(), Box<dyn std::error::Error>> {
let request = GenerateContentRequest::builder().contents(prompt).build(); let request = GenerateContentRequest::builder().contents(prompt).build();
let mut queue = gemini let mut queue = gemini
.generate_content_stream(&request, "gemini-2.0-flash-001") .stream_generate_content(&request, "gemini-2.0-flash-001")
.await?; .await?;
while let Some(Ok(response)) = queue.next().await { while let Some(response) = queue.next().await {
println!("Response: {:?}", response); println!("Response: {:?}", response?);
} }
Ok(()) Ok(())

View File

@@ -1,309 +1,218 @@
use crate::error::Result as GeminiResult; use std::vec;
use std::sync::Arc; use tracing::error;
use std::vec;
use tokio_stream::{Stream, StreamExt}; use tokio_stream::{Stream, StreamExt};
use deadqueue::unlimited::Queue; use crate::dialogue::Message;
use reqwest_eventsource::{Event, EventSource}; use crate::error::{Error, Result};
use tracing::error; use crate::network::event_source::{EventSource, ServerSentEvent};
use crate::prelude::{
use crate::dialogue::Message; Candidate, Content, CountTokensRequest, CountTokensResponse, GenerateContentRequest,
use crate::error::{Error, Result}; GenerateContentResponse, GenerateContentResponseResult, TextEmbeddingRequest,
use crate::prelude::{ TextEmbeddingResponse,
Candidate, Content, CountTokensRequest, CountTokensResponse, GenerateContentRequest, };
GenerateContentResponse, GenerateContentResponseResult, TextEmbeddingRequest, use crate::types::{PredictImageRequest, PredictImageResponse, Role};
TextEmbeddingResponse, use crate::{prelude::Part, token_provider::TokenProvider};
};
use crate::types::{PredictImageRequest, PredictImageResponse, Role}; pub const AUTH_SCOPE: &[&str] = &["https://www.googleapis.com/auth/cloud-platform"];
use crate::{prelude::Part, token_provider::TokenProvider}; const ENDPOINT_VERSION: &str = "v1beta1";
pub static AUTH_SCOPE: &[&str] = &["https://www.googleapis.com/auth/cloud-platform"]; #[derive(Clone, Debug)]
pub struct GeminiClient<T: TokenProvider + Clone> {
#[derive(Clone, Debug)] token_provider: T,
pub struct GeminiClient<T: TokenProvider + Clone> { client: reqwest::Client,
token_provider: T, api_endpoint: String,
client: reqwest::Client, project_id: String,
api_endpoint: String, location_id: String,
project_id: String, }
location_id: String,
} impl<T: TokenProvider + Clone> GeminiClient<T> {
pub fn new(
unsafe impl<T: TokenProvider + Clone> Send for GeminiClient<T> {} token_provider: T,
unsafe impl<T: TokenProvider + Clone> Sync for GeminiClient<T> {} api_endpoint: String,
project_id: String,
impl<T: TokenProvider + Clone> GeminiClient<T> { location_id: String,
pub fn new( ) -> Self {
token_provider: T, GeminiClient {
api_endpoint: String, token_provider,
project_id: String, client: reqwest::Client::new(),
location_id: String, api_endpoint,
) -> Self { project_id,
GeminiClient { location_id,
token_provider, }
client: reqwest::Client::new(), }
api_endpoint,
project_id, pub async fn generate_content(
location_id, &self,
} request: &GenerateContentRequest,
} model: &str,
) -> Result<GenerateContentResponseResult> {
pub async fn generate_content_stream( let access_token = self.token_provider.get_token(AUTH_SCOPE).await?;
&self, let endpoint_url: String = format!(
request: &GenerateContentRequest, "https://{}/{}/projects/{}/locations/{}/publishers/google/models/{}:generateContent",
model: &str, self.api_endpoint, ENDPOINT_VERSION, self.project_id, self.location_id, model,
) -> Result<impl Stream<Item = GeminiResult<GenerateContentResponseResult>>> { );
let access_token = self.token_provider.get_token(AUTH_SCOPE).await.unwrap(); let resp = self
let endpoint_url = format!( .client
"https://{}/v1beta1/projects/{}/locations/{}/publishers/google/models/{}:streamGenerateContent?alt=sse", self.api_endpoint, self.project_id, self.location_id, model, .post(&endpoint_url)
); .bearer_auth(access_token)
let client = self.client.clone(); .json(&request)
let request = request.clone(); .send()
let req = client .await?;
.post(&endpoint_url)
.bearer_auth(access_token) let txt_json = resp.text().await?;
.json(&request); tracing::debug!("generate_content response: {:?}", txt_json);
match serde_json::from_str::<GenerateContentResponse>(&txt_json) {
let event_source = EventSource::new(req).unwrap(); Ok(response) => Ok(response.into_result()?),
Err(e) => {
let mapped = event_source.filter_map(|event| { tracing::error!("Failed to parse response: {} with error {}", txt_json, e);
let event = match event { Err(e.into())
Ok(event) => event, }
Err(reqwest_eventsource::Error::StreamEnded) => { }
return Some(Err(Error::EventSourceClosedError)) }
}
Err(e) => return Some(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
let Event::Message(event_message) = event else { /// candidates. Useful for displaying incremental output as it is generated.
return None; pub async fn stream_generate_content(
}; &self,
request: &GenerateContentRequest,
let gemini_response: GenerateContentResponse = model: &str,
match serde_json::from_str(&event_message.data) { ) -> Result<impl Stream<Item = Result<GenerateContentResponseResult>> + use<'_, T>> {
Ok(gemini_response) => gemini_response, let access_token = self.token_provider.get_token(AUTH_SCOPE).await?;
Err(e) => return Some(Err(e.into())), 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 gemini_response = match gemini_response.into_result() { );
Ok(gemini_response) => gemini_response, let response = self
Err(e) => return Some(Err(e)), .client
}; .post(&endpoint_url)
.bearer_auth(access_token)
Some(Ok(gemini_response)) .json(request)
}); .send()
Ok(mapped) .await?;
} Ok(response.event_stream().filter_map(Self::parse_sse_event))
}
pub async fn stream_generate_content(
&self, fn parse_sse_event(
request: &GenerateContentRequest, event_result: std::result::Result<ServerSentEvent, tokio_util::codec::LinesCodecError>,
model: &str, ) -> Option<Result<GenerateContentResponseResult>> {
) -> Arc<Queue<Option<Result<GenerateContentResponseResult>>>> { let data = event_result.map_err(Error::from).ok()?.data?;
let queue = Arc::new(Queue::<Option<Result<GenerateContentResponseResult>>>::new()); Some(
let access_token = match self.token_provider.get_token(AUTH_SCOPE).await { serde_json::from_str::<GenerateContentResponse>(&data)
Ok(access_token) => access_token, .map_err(Error::from)
Err(e) => { .and_then(|r| r.into_result()),
queue.push(Some(Err(e))); )
return queue; }
}
}; /// Prompts a conversation to the model.
pub async fn prompt_conversation(&self, messages: &[Message], model: &str) -> Result<Message> {
// Clone the queue and other necessary data to move into the async block. let request = GenerateContentRequest {
let cloned_queue = queue.clone(); contents: messages
let endpoint_url: String = format!( .iter()
"https://{}/v1beta1/projects/{}/locations/{}/publishers/google/models/{}:streamGenerateContent?alt=sse", self.api_endpoint, self.project_id, self.location_id, model, .map(|m| Content {
); role: Some(m.role),
let client = self.client.clone(); parts: Some(vec![Part::text(m.text.clone())]),
let request = request.clone(); })
.collect(),
// Start a thread to run the request in the background. generation_config: None,
tokio::spawn(async move { tools: None,
let req = client system_instruction: None,
.post(&endpoint_url) safety_settings: None,
.bearer_auth(access_token) };
.json(&request);
let response = self.generate_content(&request, model).await?;
let mut event_source = match EventSource::new(req) {
Ok(event_source) => event_source, // Check for errors in the response.
Err(e) => { let mut candidates = GeminiClient::<T>::collect_text_from_response(&response);
cloned_queue.push(Some(Err(e.into())));
return; match candidates.pop() {
} Some(text) => Ok(Message::new(Role::Model, &text)),
}; None => Err(Error::NoCandidatesError),
while let Some(event) = event_source.next().await { }
match event { }
Ok(event) => {
if let Event::Message(event) = event { fn collect_text_from_response(response: &GenerateContentResponseResult) -> Vec<String> {
let response: serde_json::error::Result<GenerateContentResponse> = response
serde_json::from_str(&event.data); .candidates
.iter()
match response { .filter_map(Candidate::get_text)
Ok(response) => { .collect::<Vec<String>>()
let result = response.into_result(); }
let finished = match &result {
Ok(result) => result.candidates[0].finish_reason.is_some(), pub async fn text_embeddings(
Err(_) => true, &self,
}; request: &TextEmbeddingRequest,
cloned_queue.push(Some(result)); model: &str,
if finished { ) -> Result<TextEmbeddingResponse> {
break; let endpoint_url = format!(
} "https://{}/v1/projects/{}/locations/{}/publishers/google/models/{}:predict",
} self.api_endpoint, self.project_id, self.location_id, model,
Err(_) => { );
tracing::error!("Error parsing message: {}", event.data); let access_token = self.token_provider.get_token(AUTH_SCOPE).await?;
break; let resp = self
} .client
} .post(&endpoint_url)
} .bearer_auth(access_token)
} .json(&request)
Err(e) => { .send()
tracing::error!("Error in event source: {:?}", e); .await?;
break; let txt_json = resp.text().await?;
} tracing::debug!("text_embeddings response: {:?}", txt_json);
} Ok(serde_json::from_str::<TextEmbeddingResponse>(&txt_json)?)
} }
cloned_queue.push(None);
}); pub async fn count_tokens(
&self,
// Return the queue that will receive the responses. request: &CountTokensRequest,
queue model: &str,
} ) -> Result<CountTokensResponse> {
let endpoint_url = format!(
pub async fn generate_content( "https://{}/{}/projects/{}/locations/{}/publishers/google/models/{}:countTokens",
&self, self.api_endpoint, ENDPOINT_VERSION, self.project_id, self.location_id, model,
request: &GenerateContentRequest, );
model: &str, let access_token = self.token_provider.get_token(AUTH_SCOPE).await?;
) -> Result<GenerateContentResponseResult> { let resp = self
let access_token = self.token_provider.get_token(AUTH_SCOPE).await?; .client
let endpoint_url: String = format!( .post(&endpoint_url)
"https://{}/v1beta1/projects/{}/locations/{}/publishers/google/models/{}:generateContent", self.api_endpoint, self.project_id, self.location_id, model, .bearer_auth(access_token)
); .json(&request)
let resp = self .send()
.client .await?;
.post(&endpoint_url)
.bearer_auth(access_token) let txt_json = resp.text().await?;
.json(&request) tracing::debug!("count_tokens response: {:?}", txt_json);
.send() Ok(serde_json::from_str(&txt_json)?)
.await?; }
let txt_json = resp.text().await?; pub async fn predict_image(
tracing::debug!("generate_content response: {:?}", txt_json); &self,
match serde_json::from_str::<GenerateContentResponse>(&txt_json) { request: &PredictImageRequest,
Ok(response) => Ok(response.into_result()?), model: &str,
Err(e) => { ) -> Result<PredictImageResponse> {
tracing::error!("Failed to parse response: {} with error {}", txt_json, e); let endpoint_url = format!(
Err(e.into()) "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?;
/// Prompts a conversation to the model. let resp = self
pub async fn prompt_conversation(&self, messages: &[Message], model: &str) -> Result<Message> { .client
let request = GenerateContentRequest { .post(&endpoint_url)
contents: messages .bearer_auth(access_token)
.iter() .json(&request)
.map(|m| Content { .send()
role: Some(m.role), .await?;
parts: Some(vec![Part::Text(m.text.clone())]),
}) let txt_json = resp.text().await?;
.collect(),
generation_config: None, match serde_json::from_str::<PredictImageResponse>(&txt_json) {
tools: None, Ok(response) => Ok(response),
system_instruction: None, Err(e) => {
safety_settings: None, error!(response = txt_json, error = ?e, "Failed to parse response");
}; Err(e.into())
}
let response = self.generate_content(&request, model).await?; }
}
// Check for errors in the response. }
let mut candidates = GeminiClient::<T>::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<String> {
response
.candidates
.iter()
.filter_map(Candidate::get_text)
.collect::<Vec<String>>()
}
pub async fn text_embeddings(
&self,
request: &TextEmbeddingRequest,
model: &str,
) -> Result<TextEmbeddingResponse> {
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::<TextEmbeddingResponse>(&txt_json)?)
}
pub async fn count_tokens(
&self,
request: &CountTokensRequest,
model: &str,
) -> Result<CountTokensResponse> {
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<PredictImageResponse> {
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::<PredictImageResponse>(&txt_json) {
Ok(response) => Ok(response),
Err(e) => {
error!(response = txt_json, error = ?e, "Failed to parse response");
Err(e.into())
}
}
}
}

View File

@@ -1,46 +1,46 @@
use serde::{Deserialize, Serialize}; use serde::{Deserialize, Serialize};
use crate::{client::GeminiClient, error::Result, prelude::TokenProvider, types::Role}; use crate::{client::GeminiClient, error::Result, prelude::TokenProvider, types::Role};
#[derive(Clone, Debug, Serialize, Deserialize)] #[derive(Clone, Debug, Serialize, Deserialize)]
pub struct Message { pub struct Message {
pub role: Role, pub role: Role,
pub text: String, pub text: String,
} }
impl Message { impl Message {
pub fn new(role: Role, text: &str) -> Self { pub fn new(role: Role, text: &str) -> Self {
Message { Message {
role, role,
text: text.to_string(), text: text.to_string(),
} }
} }
} }
#[derive(Clone, Debug, Serialize, Deserialize)] #[derive(Clone, Debug, Serialize, Deserialize)]
pub struct Dialogue { pub struct Dialogue {
model: String, model: String,
messages: Vec<Message>, messages: Vec<Message>,
} }
impl Dialogue { impl Dialogue {
pub fn new(model: &str) -> Self { pub fn new(model: &str) -> Self {
Dialogue { Dialogue {
model: model.to_string(), model: model.to_string(),
messages: vec![], messages: vec![],
} }
} }
pub async fn do_turn<T: TokenProvider + Clone>( pub async fn do_turn<T: TokenProvider + Clone>(
&mut self, &mut self,
gemini: &GeminiClient<T>, gemini: &GeminiClient<T>,
message: &str, message: &str,
) -> Result<Message> { ) -> Result<Message> {
self.messages.push(Message::new(Role::User, message)); self.messages.push(Message::new(Role::User, message));
let response = gemini let response = gemini
.prompt_conversation(&self.messages, &self.model) .prompt_conversation(&self.messages, &self.model)
.await?; .await?;
self.messages.push(response.clone()); self.messages.push(response.clone());
Ok(response) Ok(response)
} }
} }

View File

@@ -1,90 +1,71 @@
use std::fmt::Display; use std::fmt::Display;
use reqwest_eventsource::CannotCloneRequestError; use tokio_util::codec::LinesCodecError;
use crate::types; use crate::types;
pub type Result<T> = std::result::Result<T, Error>; pub type Result<T> = std::result::Result<T, Error>;
#[derive(Debug)] #[derive(Debug)]
pub enum Error { pub enum Error {
Env(std::env::VarError), Env(std::env::VarError),
HttpClient(reqwest::Error), HttpClient(reqwest::Error),
Token(gcp_auth::Error), Token(gcp_auth::Error),
Serde(serde_json::Error), Serde(serde_json::Error),
VertexError(types::VertexApiError), VertexError(types::VertexApiError),
NoCandidatesError, NoCandidatesError,
CannotCloneRequestError(CannotCloneRequestError), /// An error occurred while decoding the SSE event stream.
EventSourceError(reqwest_eventsource::Error), EventSourceError(LinesCodecError),
EventSourceClosedError, }
}
impl Display for Error {
impl Display for Error { fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { match &self {
match &self { Error::Env(e) => write!(f, "Environment variable error: {e}"),
Error::Env(e) => write!(f, "Environment variable error: {}", e), Error::HttpClient(e) => write!(f, "HTTP Client error: {e}"),
Error::HttpClient(e) => write!(f, "HTTP Client error: {}", e), Error::Token(e) => write!(f, "Token error: {e}"),
Error::Token(e) => write!(f, "Token error: {}", e), Error::Serde(e) => write!(f, "Serde error: {e}"),
Error::Serde(e) => write!(f, "Serde error: {}", e), Error::VertexError(e) => write!(f, "Vertex error: {e}"),
Error::VertexError(e) => { Error::NoCandidatesError => write!(f, "No candidates returned for the prompt"),
write!(f, "Vertex error: {}", e) Error::EventSourceError(e) => write!(f, "EventSource error: {e}"),
} }
Error::NoCandidatesError => { }
write!(f, "No candidates returned for the prompt") }
}
Error::CannotCloneRequestError(e) => { impl std::error::Error for Error {}
write!(f, "Cannot clone request: {}", e)
} impl From<reqwest::Error> for Error {
Error::EventSourceError(e) => { fn from(e: reqwest::Error) -> Self {
write!(f, "EventSourrce Error: {}", e) Error::HttpClient(e)
} }
Error::EventSourceClosedError => { }
write!(f, "EventSource closed error")
} impl From<std::env::VarError> for Error {
} fn from(e: std::env::VarError) -> Self {
} Error::Env(e)
} }
}
impl std::error::Error for Error {}
impl From<serde_json::Error> for Error {
impl From<reqwest::Error> for Error { fn from(e: serde_json::Error) -> Self {
fn from(e: reqwest::Error) -> Self { Error::Serde(e)
Error::HttpClient(e) }
} }
}
impl From<gcp_auth::Error> for Error {
impl From<std::env::VarError> for Error { fn from(e: gcp_auth::Error) -> Self {
fn from(e: std::env::VarError) -> Self { Error::Token(e)
Error::Env(e) }
} }
}
impl From<types::VertexApiError> for Error {
impl From<gcp_auth::Error> for Error { fn from(e: types::VertexApiError) -> Self {
fn from(e: gcp_auth::Error) -> Self { Error::VertexError(e)
Error::Token(e) }
} }
}
impl From<LinesCodecError> for Error {
impl From<serde_json::Error> for Error { fn from(e: LinesCodecError) -> Self {
fn from(e: serde_json::Error) -> Self { Error::EventSourceError(e)
Error::Serde(e) }
} }
}
impl From<types::VertexApiError> for Error {
fn from(e: types::VertexApiError) -> Self {
Error::VertexError(e)
}
}
impl From<CannotCloneRequestError> for Error {
fn from(e: CannotCloneRequestError) -> Self {
Error::CannotCloneRequestError(e)
}
}
impl From<reqwest_eventsource::Error> for Error {
fn from(e: reqwest_eventsource::Error) -> Self {
Error::EventSourceError(e)
}
}

View File

@@ -1,6 +1,7 @@
mod client; mod client;
mod dialogue; mod dialogue;
pub mod error; pub mod error;
pub(crate) mod network;
mod token_provider; mod token_provider;
mod types; mod types;

108
src/network/event_source.rs Normal file
View File

@@ -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<Item = Result<ServerSentEvent, LinesCodecError>>;
}
impl EventSource for Response {
fn event_stream(self) -> impl Stream<Item = Result<ServerSentEvent, LinesCodecError>> {
stream_response(self)
}
}
/// A parsed Server-Sent Event.
#[derive(Debug, Default, Clone)]
pub struct ServerSentEvent {
pub event: Option<String>,
/// The event payload (from one or more `data:` fields, joined by `\n`).
pub data: Option<String>,
pub id: Option<String>,
pub retry: Option<usize>,
}
/// [`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<Option<Self::Item>, 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<Item = Result<ServerSentEvent, LinesCodecError>> {
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())
}

1
src/network/mod.rs Normal file
View File

@@ -0,0 +1 @@
pub mod event_source;

View File

@@ -1,18 +1,18 @@
use std::sync::Arc; use std::sync::Arc;
use crate::error::Result; use crate::error::Result;
pub trait TokenProvider { pub trait TokenProvider {
fn get_token(&self, scope: &[&str]) fn get_token(&self, scope: &[&str])
-> impl std::future::Future<Output = Result<String>> + Send; -> impl std::future::Future<Output = Result<String>> + Send;
} }
impl TokenProvider for Arc<dyn gcp_auth::TokenProvider + '_> { impl TokenProvider for Arc<dyn gcp_auth::TokenProvider + '_> {
async fn get_token(&self, scope: &[&str]) -> Result<String> { async fn get_token(&self, scope: &[&str]) -> Result<String> {
let token = self.token(scope).await; let token = self.token(scope).await;
match token { match token {
Ok(token) => Ok(token.as_str().to_string()), Ok(token) => Ok(token.as_str().to_string()),
Err(e) => Err(e.into()), Err(e) => Err(e.into()),
} }
} }
} }

View File

@@ -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::{Deserialize, Serialize};
use serde_json::Value;
#[derive(Clone, Default, Debug, Serialize, Deserialize)] #[derive(Clone, Default, Debug, Serialize, Deserialize)]
pub struct Content { pub struct Content {
@@ -13,10 +14,7 @@ impl Content {
self.parts.as_ref().map(|parts| { self.parts.as_ref().map(|parts| {
parts parts
.iter() .iter()
.filter_map(|part| match part { .filter_map(|part| part.text.clone())
Part::Text(text) => Some(text.clone()),
_ => None,
})
.collect::<String>() .collect::<String>()
}) })
} }
@@ -26,6 +24,27 @@ impl Content {
} }
} }
pub struct SystemPrompt {
content: Content,
}
impl<T: AsRef<str>> From<T> for SystemPrompt {
fn from(value: T) -> Self {
SystemPrompt {
content: Content {
parts: Some(vec![Part::text(value.as_ref())]),
role: None,
},
}
}
}
impl From<SystemPrompt> for Content {
fn from(value: SystemPrompt) -> Self {
value.content
}
}
#[derive(Default)] #[derive(Default)]
pub struct ContentBuilder { pub struct ContentBuilder {
content: Content, content: Content,
@@ -33,7 +52,7 @@ pub struct ContentBuilder {
impl ContentBuilder { impl ContentBuilder {
pub fn add_text_part<T: Into<String>>(self, text: T) -> Self { pub fn add_text_part<T: Into<String>>(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 { 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")] #[serde(rename_all = "lowercase")]
pub enum Role { pub enum Role {
User, User,
@@ -83,83 +102,83 @@ impl FromStr for Role {
} }
} }
#[derive(Clone, Debug, Serialize)] #[derive(Clone, Debug, Default, Serialize, Deserialize)]
#[serde(rename_all = "camelCase")] #[serde(rename_all = "camelCase")]
pub enum Part { pub struct Part {
Text(String), #[serde(skip_serializing_if = "Option::is_none")]
InlineData { pub text: Option<String>,
mime_type: String, #[serde(skip_serializing_if = "Option::is_none")]
data: String, pub inline_data: Option<InlineData>,
}, #[serde(skip_serializing_if = "Option::is_none")]
FileData { pub file_data: Option<FileData>,
mime_type: String, #[serde(skip_serializing_if = "Option::is_none")]
file_uri: String, pub function_call: Option<FunctionCallData>,
}, #[serde(skip_serializing_if = "Option::is_none")]
FunctionCall { pub function_response: Option<FunctionResponseData>,
name: String, /// Opaque signature emitted by Gemini 2.5+ thinking models on function-call
args: HashMap<String, String>, /// 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<String>,
/// `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<bool>,
} }
#[derive(Deserialize)] impl Part {
#[serde(rename_all = "camelCase")] pub fn text(text: impl Into<String>) -> Self {
struct PartHelper { Part {
text: Option<String>, text: Some(text.into()),
inline_data: Option<InlineDataHelper>, ..Default::default()
file_data: Option<FileDataHelper>,
function_call: Option<FunctionCallHelper>,
}
#[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<String, String>,
}
impl<'de> Deserialize<'de> for Part {
fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
where
D: serde::Deserializer<'de>,
{
let helper = PartHelper::deserialize(deserializer)?;
if let Some(text) = helper.text {
return Ok(Part::Text(text));
} }
if let Some(inline_data) = helper.inline_data { }
return Ok(Part::InlineData {
mime_type: inline_data.mime_type, pub fn function_call(name: impl Into<String>, args: Value) -> Self {
data: inline_data.data, 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, pub fn function_response(name: impl Into<String>, response: Value) -> Self {
file_uri: file_data.file_uri, 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,
}

View File

@@ -22,7 +22,7 @@ impl CountTokensRequestBuilder {
pub fn from_prompt(prompt: &str) -> Self { pub fn from_prompt(prompt: &str) -> Self {
CountTokensRequestBuilder { CountTokensRequestBuilder {
contents: Content { contents: Content {
parts: Some(vec![super::Part::Text(prompt.to_string())]), parts: Some(vec![super::Part::text(prompt)]),
..Default::default() ..Default::default()
}, },
} }

View File

@@ -1,5 +1,3 @@
use std::collections::HashMap;
use serde::{Deserialize, Serialize}; use serde::{Deserialize, Serialize};
use serde_json::Value; use serde_json::Value;
@@ -125,6 +123,37 @@ pub struct GenerationConfig {
pub response_mime_type: Option<String>, pub response_mime_type: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")] #[serde(skip_serializing_if = "Option::is_none")]
pub response_schema: Option<Value>, pub response_schema: Option<Value>,
#[serde(skip_serializing_if = "Option::is_none")]
pub thinking_config: Option<ThinkingConfig>,
}
/// Configures Gemini 2.5+ thinking models. See
/// <https://docs.cloud.google.com/vertex-ai/generative-ai/docs/thinking>.
#[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<bool>,
/// 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<i32>,
/// Coarse-grained thinking effort. Mutually-exclusive shorthand for
/// `thinking_budget` on supported models.
#[serde(skip_serializing_if = "Option::is_none")]
pub thinking_level: Option<ThinkingLevel>,
}
/// 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 { impl GenerationConfig {
@@ -184,6 +213,11 @@ impl GenerationConfigBuilder {
self 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 { pub fn build(self) -> GenerationConfig {
self.generation_config self.generation_config
} }
@@ -294,22 +328,7 @@ pub struct UsageMetadata {
pub struct FunctionDeclaration { pub struct FunctionDeclaration {
pub name: String, pub name: String,
pub description: String, pub description: String,
pub parameters: FunctionParameters, pub parameters: Value,
}
#[derive(Clone, Debug, Serialize, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct FunctionParameters {
pub r#type: String,
pub properties: HashMap<String, FunctionParametersProperty>,
pub required: Vec<String>,
}
#[derive(Clone, Debug, Serialize, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct FunctionParametersProperty {
pub r#type: String,
pub description: String,
} }
#[derive(Debug, Serialize, Deserialize)] #[derive(Debug, Serialize, Deserialize)]
@@ -485,6 +504,47 @@ mod tests {
serde_json::from_str::<GenerateContentResponse>(input).unwrap(); serde_json::from_str::<GenerateContentResponse>(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::<GenerateContentResponse>(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] #[test]
fn parses_safety_rating_without_scores() { fn parses_safety_rating_without_scores() {
let input = r#"{ let input = r#"{
@@ -531,34 +591,4 @@ mod tests {
}"#; }"#;
serde_json::from_str::<GenerateContentResponse>(input).unwrap(); serde_json::from_str::<GenerateContentResponse>(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::<GenerateContentResponse>(input).unwrap();
}
} }

View File

@@ -12,6 +12,10 @@ pub struct PredictImageRequest {
pub struct PredictImageRequestPrompt { pub struct PredictImageRequestPrompt {
/// The text prompt for the image. /// The text prompt for the image.
/// The following models support different values for this parameter: /// 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-generate-001`: up to 480 tokens.
/// - `imagen-3.0-fast-generate-001`: up to 480 tokens. /// - `imagen-3.0-fast-generate-001`: up to 480 tokens.
/// - `imagegeneration@006`: up to 128 tokens. /// - `imagegeneration@006`: up to 128 tokens.
@@ -24,27 +28,20 @@ pub struct PredictImageRequestPrompt {
#[serde(rename_all = "camelCase")] #[serde(rename_all = "camelCase")]
pub struct PredictImageRequestParameters { pub struct PredictImageRequestParameters {
/// The number of images to generate. The default value is 4. /// The number of images to generate. The default value is 4.
/// The following models support different values for this parameter: /// Supported range: 1 to 4 for all current models.
/// - `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.
pub sample_count: i32, pub sample_count: i32,
/// The random seed for image generation. This is not available when addWatermark is set to /// The random seed for image generation. This is not available when `add_watermark` is set to
/// true. /// `true`.
#[serde(skip_serializing_if = "Option::is_none")] #[serde(skip_serializing_if = "Option::is_none")]
pub seed: Option<u32>, pub seed: Option<u32>,
/// 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<bool>,
/// A description of what to discourage in the generated images. /// A description of what to discourage in the generated images.
/// The following models support this parameter: /// 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-generate-001`: up to 480 tokens.
/// - `imagen-3.0-fast-generate-001`: up to 480 tokens. /// - `imagen-3.0-fast-generate-001`: up to 480 tokens.
/// - `imagegeneration@006`: up to 128 tokens. /// - `imagegeneration@006`: up to 128 tokens.
@@ -53,13 +50,8 @@ pub struct PredictImageRequestParameters {
#[serde(skip_serializing_if = "Option::is_none")] #[serde(skip_serializing_if = "Option::is_none")]
pub negative_prompt: Option<String>, pub negative_prompt: Option<String>,
/// The aspect ratio for the image. The default value is "1:1". /// The aspect ratio for the image. The default value is `"1:1"`.
/// The following models support different values for this parameter: /// Supported values: `"1:1"`, `"9:16"`, `"16:9"`, `"3:4"`, `"4:3"`.
/// - `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".
#[serde(skip_serializing_if = "Option::is_none")] #[serde(skip_serializing_if = "Option::is_none")]
pub aspect_ratio: Option<String>, pub aspect_ratio: Option<String>,
@@ -77,7 +69,7 @@ pub struct PredictImageRequestParameters {
/// - "cyberpunk" /// - "cyberpunk"
/// - "pop_art" /// - "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")] #[serde(skip_serializing_if = "Option::is_none")]
pub sample_image_style: Option<String>, pub sample_image_style: Option<String>,
@@ -87,56 +79,41 @@ pub struct PredictImageRequestParameters {
/// - `"allow_all"`: Allow generation of people of all ages. /// - `"allow_all"`: Allow generation of people of all ages.
/// ///
/// The default value is `"allow_adult"`. /// 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")] #[serde(skip_serializing_if = "Option::is_none")]
pub person_generation: Option<PersonGeneration>, pub person_generation: Option<String>,
/// 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<String>,
/// Adds a filter level to safety filtering. The following values are supported: /// Adds a filter level to safety filtering. The following values are supported:
/// /// - `"block_low_and_above"`: Strongest filtering, most strict blocking.
/// - "block_low_and_above": Strongest filtering level, most strict blocking. /// - `"block_medium_and_above"`: Block medium and high severity content. Default.
/// Deprecated value: "block_most". /// - `"block_only_high"`: Only block high severity content.
/// - "block_medium_and_above": Block some problematic prompts and responses. /// - `"block_none"`: No blocking.
/// 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.
#[serde(skip_serializing_if = "Option::is_none")] #[serde(skip_serializing_if = "Option::is_none")]
pub safety_setting: Option<PredictImageSafetySetting>, pub safety_setting: Option<String>,
/// Add an invisible watermark to the generated images. The default value is `false` for the /// Add an invisible SynthID watermark to the generated images. Defaults to `true` for
/// `imagegeneration@002` and `imagegeneration@005` models, and `true` for the /// Imagen 3 and 4 models. Not available when `seed` is set.
/// `imagen-3.0-fast-generate-001`, `imagegeneration@006`, and imagegeneration@006 models.
#[serde(skip_serializing_if = "Option::is_none")] #[serde(skip_serializing_if = "Option::is_none")]
pub add_watermark: Option<bool>, pub add_watermark: Option<bool>,
/// Cloud Storage URI to store the generated images. /// Cloud Storage URI to store the generated images.
#[serde(skip_serializing_if = "Option::is_none")] #[serde(skip_serializing_if = "Option::is_none")]
pub storage_uri: Option<String>, pub storage_uri: Option<String>,
/// 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<bool>,
/// 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<String>,
/// Output image resolution. `"1K"` (default) or `"2K"`.
#[serde(skip_serializing_if = "Option::is_none")]
pub sample_image_size: Option<String>,
} }
#[derive(Debug, Serialize, Deserialize)] #[derive(Debug, Serialize, Deserialize)]
@@ -144,10 +121,9 @@ pub struct PredictImageRequestParameters {
pub struct PredictImageRequestParametersOutputOptions { pub struct PredictImageRequestParametersOutputOptions {
/// The image format that the output should be saved as. The following values are supported: /// The image format that the output should be saved as. The following values are supported:
/// ///
/// - "image/png": Save as a PNG image /// - `"image/png"`: Save as a PNG image (default)
/// - "image/jpeg": Save as a JPEG image /// - `"image/jpeg"`: Save as a JPEG image
/// /// - `"image/webp"`: Save as a WebP image
/// The default value is "image/png".v
pub mime_type: Option<String>, pub mime_type: Option<String>,
/// The level of compression if the output type is "image/jpeg". /// The level of compression if the output type is "image/jpeg".
@@ -168,20 +144,3 @@ pub struct PredictImageResponsePrediction {
pub bytes_base64_encoded: Vec<u8>, pub bytes_base64_encoded: Vec<u8>,
pub mime_type: String, 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,
}