feat: add SSE streaming, thinking models, and modernized Part struct
Some checks failed
Rust / build (push) Has been cancelled
Some checks failed
Rust / build (push) Has been cancelled
This commit is contained in:
24
Cargo.toml
24
Cargo.toml
@@ -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"
|
||||||
|
|||||||
@@ -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()
|
||||||
},
|
},
|
||||||
};
|
};
|
||||||
|
|||||||
@@ -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()),
|
||||||
|
|||||||
@@ -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()),
|
||||||
|
|||||||
@@ -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()),
|
||||||
|
|||||||
@@ -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,
|
||||||
|
|||||||
@@ -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(())
|
||||||
|
|||||||
181
src/client.rs
181
src/client.rs
@@ -1,14 +1,11 @@
|
|||||||
use crate::error::Result as GeminiResult;
|
|
||||||
use std::sync::Arc;
|
|
||||||
use std::vec;
|
use std::vec;
|
||||||
use tokio_stream::{Stream, StreamExt};
|
|
||||||
|
|
||||||
use deadqueue::unlimited::Queue;
|
|
||||||
use reqwest_eventsource::{Event, EventSource};
|
|
||||||
use tracing::error;
|
use tracing::error;
|
||||||
|
|
||||||
|
use tokio_stream::{Stream, StreamExt};
|
||||||
|
|
||||||
use crate::dialogue::Message;
|
use crate::dialogue::Message;
|
||||||
use crate::error::{Error, Result};
|
use crate::error::{Error, Result};
|
||||||
|
use crate::network::event_source::{EventSource, ServerSentEvent};
|
||||||
use crate::prelude::{
|
use crate::prelude::{
|
||||||
Candidate, Content, CountTokensRequest, CountTokensResponse, GenerateContentRequest,
|
Candidate, Content, CountTokensRequest, CountTokensResponse, GenerateContentRequest,
|
||||||
GenerateContentResponse, GenerateContentResponseResult, TextEmbeddingRequest,
|
GenerateContentResponse, GenerateContentResponseResult, TextEmbeddingRequest,
|
||||||
@@ -17,7 +14,8 @@ use crate::prelude::{
|
|||||||
use crate::types::{PredictImageRequest, PredictImageResponse, Role};
|
use crate::types::{PredictImageRequest, PredictImageResponse, Role};
|
||||||
use crate::{prelude::Part, token_provider::TokenProvider};
|
use crate::{prelude::Part, token_provider::TokenProvider};
|
||||||
|
|
||||||
pub static AUTH_SCOPE: &[&str] = &["https://www.googleapis.com/auth/cloud-platform"];
|
pub const AUTH_SCOPE: &[&str] = &["https://www.googleapis.com/auth/cloud-platform"];
|
||||||
|
const ENDPOINT_VERSION: &str = "v1beta1";
|
||||||
|
|
||||||
#[derive(Clone, Debug)]
|
#[derive(Clone, Debug)]
|
||||||
pub struct GeminiClient<T: TokenProvider + Clone> {
|
pub struct GeminiClient<T: TokenProvider + Clone> {
|
||||||
@@ -28,9 +26,6 @@ pub struct GeminiClient<T: TokenProvider + Clone> {
|
|||||||
location_id: String,
|
location_id: String,
|
||||||
}
|
}
|
||||||
|
|
||||||
unsafe impl<T: TokenProvider + Clone> Send for GeminiClient<T> {}
|
|
||||||
unsafe impl<T: TokenProvider + Clone> Sync for GeminiClient<T> {}
|
|
||||||
|
|
||||||
impl<T: TokenProvider + Clone> GeminiClient<T> {
|
impl<T: TokenProvider + Clone> GeminiClient<T> {
|
||||||
pub fn new(
|
pub fn new(
|
||||||
token_provider: T,
|
token_provider: T,
|
||||||
@@ -47,128 +42,6 @@ impl<T: TokenProvider + Clone> GeminiClient<T> {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
pub async fn generate_content_stream(
|
|
||||||
&self,
|
|
||||||
request: &GenerateContentRequest,
|
|
||||||
model: &str,
|
|
||||||
) -> Result<impl Stream<Item = GeminiResult<GenerateContentResponseResult>>> {
|
|
||||||
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<Queue<Option<Result<GenerateContentResponseResult>>>> {
|
|
||||||
let queue = Arc::new(Queue::<Option<Result<GenerateContentResponseResult>>>::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<GenerateContentResponse> =
|
|
||||||
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(
|
pub async fn generate_content(
|
||||||
&self,
|
&self,
|
||||||
request: &GenerateContentRequest,
|
request: &GenerateContentRequest,
|
||||||
@@ -176,7 +49,8 @@ impl<T: TokenProvider + Clone> GeminiClient<T> {
|
|||||||
) -> Result<GenerateContentResponseResult> {
|
) -> Result<GenerateContentResponseResult> {
|
||||||
let access_token = self.token_provider.get_token(AUTH_SCOPE).await?;
|
let access_token = self.token_provider.get_token(AUTH_SCOPE).await?;
|
||||||
let endpoint_url: String = format!(
|
let endpoint_url: String = format!(
|
||||||
"https://{}/v1beta1/projects/{}/locations/{}/publishers/google/models/{}:generateContent", self.api_endpoint, self.project_id, self.location_id, model,
|
"https://{}/{}/projects/{}/locations/{}/publishers/google/models/{}:generateContent",
|
||||||
|
self.api_endpoint, ENDPOINT_VERSION, self.project_id, self.location_id, model,
|
||||||
);
|
);
|
||||||
let resp = self
|
let resp = self
|
||||||
.client
|
.client
|
||||||
@@ -197,6 +71,41 @@ impl<T: TokenProvider + Clone> GeminiClient<T> {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
/// 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<impl Stream<Item = Result<GenerateContentResponseResult>> + 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<ServerSentEvent, tokio_util::codec::LinesCodecError>,
|
||||||
|
) -> Option<Result<GenerateContentResponseResult>> {
|
||||||
|
let data = event_result.map_err(Error::from).ok()?.data?;
|
||||||
|
Some(
|
||||||
|
serde_json::from_str::<GenerateContentResponse>(&data)
|
||||||
|
.map_err(Error::from)
|
||||||
|
.and_then(|r| r.into_result()),
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
/// Prompts a conversation to the model.
|
/// Prompts a conversation to the model.
|
||||||
pub async fn prompt_conversation(&self, messages: &[Message], model: &str) -> Result<Message> {
|
pub async fn prompt_conversation(&self, messages: &[Message], model: &str) -> Result<Message> {
|
||||||
let request = GenerateContentRequest {
|
let request = GenerateContentRequest {
|
||||||
@@ -204,7 +113,7 @@ impl<T: TokenProvider + Clone> GeminiClient<T> {
|
|||||||
.iter()
|
.iter()
|
||||||
.map(|m| Content {
|
.map(|m| Content {
|
||||||
role: Some(m.role),
|
role: Some(m.role),
|
||||||
parts: Some(vec![Part::Text(m.text.clone())]),
|
parts: Some(vec![Part::text(m.text.clone())]),
|
||||||
})
|
})
|
||||||
.collect(),
|
.collect(),
|
||||||
generation_config: None,
|
generation_config: None,
|
||||||
@@ -260,8 +169,8 @@ impl<T: TokenProvider + Clone> GeminiClient<T> {
|
|||||||
model: &str,
|
model: &str,
|
||||||
) -> Result<CountTokensResponse> {
|
) -> Result<CountTokensResponse> {
|
||||||
let endpoint_url = format!(
|
let endpoint_url = format!(
|
||||||
"https://{}/v1beta1/projects/{}/locations/{}/publishers/google/models/{}:countTokens",
|
"https://{}/{}/projects/{}/locations/{}/publishers/google/models/{}:countTokens",
|
||||||
self.api_endpoint, self.project_id, self.location_id, model,
|
self.api_endpoint, ENDPOINT_VERSION, self.project_id, self.location_id, model,
|
||||||
);
|
);
|
||||||
let access_token = self.token_provider.get_token(AUTH_SCOPE).await?;
|
let access_token = self.token_provider.get_token(AUTH_SCOPE).await?;
|
||||||
let resp = self
|
let resp = self
|
||||||
|
|||||||
55
src/error.rs
55
src/error.rs
@@ -1,6 +1,6 @@
|
|||||||
use std::fmt::Display;
|
use std::fmt::Display;
|
||||||
|
|
||||||
use reqwest_eventsource::CannotCloneRequestError;
|
use tokio_util::codec::LinesCodecError;
|
||||||
|
|
||||||
use crate::types;
|
use crate::types;
|
||||||
|
|
||||||
@@ -14,33 +14,20 @@ pub enum 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) => {
|
Error::VertexError(e) => write!(f, "Vertex error: {e}"),
|
||||||
write!(f, "Vertex error: {}", e)
|
Error::NoCandidatesError => write!(f, "No candidates returned for the prompt"),
|
||||||
}
|
Error::EventSourceError(e) => write!(f, "EventSource 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")
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -59,32 +46,26 @@ impl From<std::env::VarError> for Error {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
impl From<gcp_auth::Error> for Error {
|
|
||||||
fn from(e: gcp_auth::Error) -> Self {
|
|
||||||
Error::Token(e)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
impl From<serde_json::Error> for Error {
|
impl From<serde_json::Error> for Error {
|
||||||
fn from(e: serde_json::Error) -> Self {
|
fn from(e: serde_json::Error) -> Self {
|
||||||
Error::Serde(e)
|
Error::Serde(e)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
impl From<gcp_auth::Error> for Error {
|
||||||
|
fn from(e: gcp_auth::Error) -> Self {
|
||||||
|
Error::Token(e)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
impl From<types::VertexApiError> for Error {
|
impl From<types::VertexApiError> for Error {
|
||||||
fn from(e: types::VertexApiError) -> Self {
|
fn from(e: types::VertexApiError) -> Self {
|
||||||
Error::VertexError(e)
|
Error::VertexError(e)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
impl From<CannotCloneRequestError> for Error {
|
impl From<LinesCodecError> for Error {
|
||||||
fn from(e: CannotCloneRequestError) -> Self {
|
fn from(e: LinesCodecError) -> Self {
|
||||||
Error::CannotCloneRequestError(e)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
impl From<reqwest_eventsource::Error> for Error {
|
|
||||||
fn from(e: reqwest_eventsource::Error) -> Self {
|
|
||||||
Error::EventSourceError(e)
|
Error::EventSourceError(e)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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
108
src/network/event_source.rs
Normal 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
1
src/network/mod.rs
Normal file
@@ -0,0 +1 @@
|
|||||||
|
pub mod event_source;
|
||||||
@@ -4,7 +4,7 @@ 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 + '_> {
|
||||||
|
|||||||
@@ -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,
|
||||||
|
}
|
||||||
|
|||||||
@@ -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()
|
||||||
},
|
},
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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();
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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,
|
|
||||||
}
|
|
||||||
|
|||||||
Reference in New Issue
Block a user