refactor: extract google ai logic to windmill-common and use native gemini api in chat proxy (#8115)
* refactor: extract google ai logic to windmill-common and use native gemini api in chat proxy
Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
* fix: use x-goog-api-key header for google ai non-chat requests
Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
* fix: transform gemini models response to openai format and use correct auth header
Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
* fix: skip thought parts from gemini thinking models in sse stream
Co-Authored-By: Claude Opus 4.5 <noreply@anthropic.com>
* Revert "fix: skip thought parts from gemini thinking models in sse stream"
This reverts commit dfa01d282c.
* fix: handle tool calls and sanitize schemas in gemini chat proxy
Co-Authored-By: Claude Opus 4.5 <noreply@anthropic.com>
* refactor: move Gemini→OpenAI response conversion to windmill-common
Extract streaming and non-streaming Gemini response conversion into
shared functions in ai_google so the API proxy and worker use the same
logic instead of duplicating format translation.
Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
* fix: review fixes for google ai refactor
- Remove duplicate parse_data_url from worker utils, use shared version
from windmill_common::ai_google in both google_ai and anthropic providers
- Improve error diagnostics in google.rs by including HTTP status code
in error messages from Gemini API responses
- Change GeminiToolCallEvent::into_extra_content to instance method
to_extra_content using &self
Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
* refactor: deduplicate worker Gemini message conversion using pre-flight pattern
Replace the worker's `convert_messages_to_gemini` and
`convert_content_to_parts_with_s3` (~130 lines) with the existing
pre-flight pattern: `prepare_messages_for_api` converts S3 objects to
data URLs, then the shared `openai_messages_to_gemini` handles the rest.
Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
---------
Co-authored-by: Claude Sonnet 4.6 <noreply@anthropic.com>
Co-authored-by: hugocasa <hugo@casademont.ch>
This commit is contained in:
1
backend/Cargo.lock
generated
1
backend/Cargo.lock
generated
@@ -15849,6 +15849,7 @@ dependencies = [
|
||||
"dashmap 6.1.0",
|
||||
"datafusion",
|
||||
"ed25519-dalek",
|
||||
"eventsource-stream",
|
||||
"flate2",
|
||||
"futures",
|
||||
"git-version",
|
||||
|
||||
@@ -174,6 +174,7 @@ aws-sdk-bedrock = { workspace = true, optional = true }
|
||||
aws-sdk-bedrockruntime = { workspace = true, optional = true }
|
||||
aws-smithy-types = { workspace = true, optional = true }
|
||||
async-trait.workspace = true
|
||||
eventsource-stream.workspace = true
|
||||
windmill-jseval.workspace = true
|
||||
tar.workspace = true
|
||||
flate2.workspace = true
|
||||
|
||||
@@ -29,7 +29,7 @@ const AI_TIMEOUT_MAX_SECS: u64 = 86400; // 24 hours
|
||||
const AI_TIMEOUT_DEFAULT_SECS: u64 = 3600; // 1 hour
|
||||
const HTTP_POOL_MAX_IDLE_PER_HOST: usize = 10;
|
||||
const HTTP_POOL_IDLE_TIMEOUT_SECS: u64 = 90;
|
||||
const KEEPALIVE_INTERVAL_SECS: u64 = 15;
|
||||
pub(crate) const KEEPALIVE_INTERVAL_SECS: u64 = 15;
|
||||
|
||||
lazy_static::lazy_static! {
|
||||
/// AI request timeout in seconds.
|
||||
@@ -87,7 +87,7 @@ lazy_static::lazy_static! {
|
||||
}
|
||||
};
|
||||
|
||||
static ref HTTP_CLIENT: Client = configure_client(reqwest::ClientBuilder::new()
|
||||
pub(crate) static ref HTTP_CLIENT: Client = configure_client(reqwest::ClientBuilder::new()
|
||||
.timeout(std::time::Duration::from_secs(*AI_TIMEOUT_SECS))
|
||||
.pool_max_idle_per_host(HTTP_POOL_MAX_IDLE_PER_HOST)
|
||||
.pool_idle_timeout(Some(std::time::Duration::from_secs(HTTP_POOL_IDLE_TIMEOUT_SECS)))
|
||||
@@ -378,12 +378,7 @@ impl AIRequestConfig {
|
||||
let is_anthropic_sdk = headers.get("X-Anthropic-SDK").is_some();
|
||||
let is_google_ai = matches!(provider, AIProvider::GoogleAI);
|
||||
|
||||
// GoogleAI uses OpenAI-compatible endpoint in the proxy (for the chat), but not for the ai agent
|
||||
let base_url = if is_google_ai {
|
||||
format!("{}/openai", base_url)
|
||||
} else {
|
||||
base_url.to_string()
|
||||
};
|
||||
let base_url = base_url.to_string();
|
||||
let base_url = base_url.as_str();
|
||||
|
||||
// Build URL based on provider
|
||||
@@ -428,6 +423,9 @@ impl AIRequestConfig {
|
||||
if let Some(api_key) = self.api_key {
|
||||
if is_azure {
|
||||
request = request.header("api-key", api_key.clone())
|
||||
} else if is_google_ai {
|
||||
// Native Gemini API uses x-goog-api-key, not Authorization: Bearer
|
||||
request = request.header("x-goog-api-key", api_key.clone())
|
||||
} else {
|
||||
request = request.header("authorization", format!("Bearer {}", api_key.clone()))
|
||||
}
|
||||
@@ -611,7 +609,7 @@ fn is_sse_response(headers: &HeaderMap) -> bool {
|
||||
.unwrap_or(false)
|
||||
}
|
||||
|
||||
fn inject_keepalives<S>(
|
||||
pub(crate) fn inject_keepalives<S>(
|
||||
upstream: S,
|
||||
interval: Duration,
|
||||
) -> impl futures::Stream<Item = std::result::Result<Bytes, reqwest::Error>>
|
||||
@@ -830,6 +828,36 @@ async fn proxy(
|
||||
ai_path = chat_path;
|
||||
}
|
||||
|
||||
// Handle GoogleAI (Gemini) using the native Gemini API
|
||||
if matches!(provider, AIProvider::GoogleAI) {
|
||||
let api_key = request_config.api_key.as_deref().unwrap_or("");
|
||||
let base_url = request_config.base_url.trim_end_matches('/');
|
||||
|
||||
let mut tx = db.begin().await?;
|
||||
audit_log(
|
||||
&mut *tx,
|
||||
&authed,
|
||||
"ai.request",
|
||||
ActionKind::Execute,
|
||||
&w_id,
|
||||
Some(&authed.email),
|
||||
Some([("ai_config_path", &format!("{:?}", ai_path)[..])].into()),
|
||||
)
|
||||
.await?;
|
||||
tx.commit().await?;
|
||||
|
||||
return match ai_path.as_str() {
|
||||
"chat/completions" => {
|
||||
crate::google::handle_google_ai_chat(&body, api_key, base_url).await
|
||||
}
|
||||
"models" => crate::google::handle_google_ai_models(api_key, base_url).await,
|
||||
_ => Err(Error::BadRequest(format!(
|
||||
"Unsupported Google AI path: {}",
|
||||
ai_path
|
||||
))),
|
||||
};
|
||||
}
|
||||
|
||||
// Handle Bedrock-specific logic when the feature is enabled
|
||||
#[cfg(feature = "bedrock")]
|
||||
{
|
||||
|
||||
306
backend/windmill-api/src/google.rs
Normal file
306
backend/windmill-api/src/google.rs
Normal file
@@ -0,0 +1,306 @@
|
||||
//! Google AI (Gemini API) handler for the AI chat proxy.
|
||||
//!
|
||||
//! Handles POST `chat/completions` requests using the native Gemini API,
|
||||
//! converting from/to OpenAI format so the existing frontend parsers continue to work.
|
||||
//!
|
||||
//! Used by `windmill-api/src/ai.rs` when the provider is `GoogleAI`.
|
||||
//! Shared conversion logic lives in `windmill_common::ai_google`.
|
||||
|
||||
use axum::body::Body;
|
||||
use bytes::Bytes;
|
||||
use eventsource_stream::Eventsource;
|
||||
use futures::StreamExt;
|
||||
use serde::Deserialize;
|
||||
use serde_json::json;
|
||||
use windmill_common::{
|
||||
ai_google::{
|
||||
gemini_event_to_openai_sse_chunks, gemini_response_to_openai, openai_messages_to_gemini,
|
||||
parse_gemini_response, parse_gemini_sse_event, sanitize_schema_for_google,
|
||||
GeminiFunctionDeclaration, GeminiGenerationConfig, GeminiTextRequest, GeminiTool,
|
||||
},
|
||||
ai_types::OpenAIMessage,
|
||||
error::{Error, Result},
|
||||
};
|
||||
|
||||
use crate::ai::{inject_keepalives, HTTP_CLIENT, KEEPALIVE_INTERVAL_SECS};
|
||||
|
||||
// ============================================================================
|
||||
// Request type (OpenAI format received from the frontend)
|
||||
// ============================================================================
|
||||
|
||||
#[derive(Deserialize, Debug)]
|
||||
struct ChatRequest {
|
||||
model: String,
|
||||
messages: Vec<OpenAIMessage>,
|
||||
#[serde(default)]
|
||||
stream: bool,
|
||||
#[serde(default)]
|
||||
temperature: Option<f32>,
|
||||
#[serde(default)]
|
||||
max_tokens: Option<u32>,
|
||||
#[serde(default)]
|
||||
tools: Option<Vec<ChatRequestTool>>,
|
||||
}
|
||||
|
||||
#[derive(Deserialize, Debug)]
|
||||
struct ChatRequestTool {
|
||||
function: ChatRequestToolFunction,
|
||||
}
|
||||
|
||||
#[derive(Deserialize, Debug)]
|
||||
struct ChatRequestToolFunction {
|
||||
name: String,
|
||||
#[serde(default)]
|
||||
description: Option<String>,
|
||||
#[serde(default)]
|
||||
parameters: Option<serde_json::Value>,
|
||||
}
|
||||
|
||||
// ============================================================================
|
||||
// Public handler
|
||||
// ============================================================================
|
||||
|
||||
/// Handle a `chat/completions` POST request using the native Gemini API.
|
||||
///
|
||||
/// Converts the incoming OpenAI-format body to a `GeminiTextRequest`, sends it
|
||||
/// to the appropriate Gemini endpoint, and converts the response back to the
|
||||
/// OpenAI SSE or JSON format that the frontend expects.
|
||||
pub async fn handle_google_ai_chat(
|
||||
body: &Bytes,
|
||||
api_key: &str,
|
||||
base_url: &str,
|
||||
) -> Result<(http::StatusCode, http::HeaderMap, Body)> {
|
||||
let request: ChatRequest = serde_json::from_slice(body)
|
||||
.map_err(|e| Error::BadRequest(format!("Failed to parse request body: {}", e)))?;
|
||||
|
||||
let (contents, system_instruction) = openai_messages_to_gemini(&request.messages);
|
||||
|
||||
let generation_config =
|
||||
if request.temperature.is_some() || request.max_tokens.is_some() {
|
||||
Some(GeminiGenerationConfig {
|
||||
temperature: request.temperature,
|
||||
max_output_tokens: request.max_tokens,
|
||||
response_mime_type: None,
|
||||
response_schema: None,
|
||||
})
|
||||
} else {
|
||||
None
|
||||
};
|
||||
|
||||
let gemini_tools = request.tools.as_ref().map(|tools| {
|
||||
let declarations: Vec<GeminiFunctionDeclaration> = tools
|
||||
.iter()
|
||||
.map(|t| {
|
||||
let mut params = t.function.parameters.clone().unwrap_or(json!({}));
|
||||
sanitize_schema_for_google(&mut params);
|
||||
GeminiFunctionDeclaration {
|
||||
name: t.function.name.clone(),
|
||||
description: t.function.description.clone(),
|
||||
parameters: params,
|
||||
}
|
||||
})
|
||||
.collect();
|
||||
vec![GeminiTool {
|
||||
function_declarations: Some(declarations),
|
||||
google_search: None,
|
||||
}]
|
||||
});
|
||||
|
||||
let gemini_request = GeminiTextRequest {
|
||||
contents,
|
||||
tools: gemini_tools,
|
||||
tool_config: None,
|
||||
system_instruction,
|
||||
generation_config,
|
||||
};
|
||||
|
||||
let request_body = serde_json::to_string(&gemini_request)
|
||||
.map_err(|e| Error::internal_err(format!("Failed to serialize Gemini request: {}", e)))?;
|
||||
|
||||
let base_url = base_url.trim_end_matches('/');
|
||||
|
||||
if request.stream {
|
||||
handle_streaming(&request.model, request_body, api_key, base_url).await
|
||||
} else {
|
||||
handle_non_streaming(&request.model, request_body, api_key, base_url).await
|
||||
}
|
||||
}
|
||||
|
||||
// ============================================================================
|
||||
// Streaming path
|
||||
// ============================================================================
|
||||
|
||||
async fn handle_streaming(
|
||||
model: &str,
|
||||
request_body: String,
|
||||
api_key: &str,
|
||||
base_url: &str,
|
||||
) -> Result<(http::StatusCode, http::HeaderMap, Body)> {
|
||||
let endpoint = format!("{}/models/{}:streamGenerateContent?alt=sse", base_url, model);
|
||||
|
||||
let response = HTTP_CLIENT
|
||||
.post(&endpoint)
|
||||
.header("content-type", "application/json")
|
||||
.header("x-goog-api-key", api_key)
|
||||
.body(request_body)
|
||||
.send()
|
||||
.await
|
||||
.map_err(|e| {
|
||||
Error::internal_err(format!("Failed to send request to Gemini API: {}", e))
|
||||
})?;
|
||||
|
||||
if let Err(e) = response.error_for_status_ref() {
|
||||
let status = e.status().map(|s| s.to_string()).unwrap_or_default();
|
||||
let body = response.text().await.unwrap_or_default();
|
||||
return Err(Error::AIError(format!("{}: {}", status, body)));
|
||||
}
|
||||
|
||||
let id = format!("chatcmpl-{}", uuid::Uuid::new_v4().simple());
|
||||
let model_str = model.to_string();
|
||||
|
||||
let gemini_sse_stream = response.bytes_stream().eventsource();
|
||||
let openai_sse_stream = async_stream::stream! {
|
||||
tokio::pin!(gemini_sse_stream);
|
||||
let mut tool_call_index: usize = 0;
|
||||
while let Some(event) = gemini_sse_stream.next().await {
|
||||
match event {
|
||||
Ok(event) => match parse_gemini_sse_event(&event.data) {
|
||||
Ok(Some(parsed)) => {
|
||||
for chunk in gemini_event_to_openai_sse_chunks(
|
||||
&parsed, &id, &model_str, &mut tool_call_index,
|
||||
) {
|
||||
yield Ok::<Bytes, reqwest::Error>(Bytes::from(chunk));
|
||||
}
|
||||
}
|
||||
Ok(None) => {}
|
||||
Err(e) => tracing::error!("Error parsing Gemini SSE event: {}", e),
|
||||
},
|
||||
Err(e) => tracing::error!("Error reading Gemini SSE stream: {}", e),
|
||||
}
|
||||
}
|
||||
yield Ok::<Bytes, reqwest::Error>(Bytes::from("data: [DONE]\n\n"));
|
||||
};
|
||||
|
||||
let mut headers = http::HeaderMap::new();
|
||||
headers.insert("content-type", "text/event-stream".parse().unwrap());
|
||||
headers.insert("cache-control", "no-cache".parse().unwrap());
|
||||
headers.insert("connection", "keep-alive".parse().unwrap());
|
||||
|
||||
Ok((
|
||||
http::StatusCode::OK,
|
||||
headers,
|
||||
Body::from_stream(inject_keepalives(
|
||||
Box::pin(openai_sse_stream),
|
||||
std::time::Duration::from_secs(KEEPALIVE_INTERVAL_SECS),
|
||||
)),
|
||||
))
|
||||
}
|
||||
|
||||
// ============================================================================
|
||||
// Model listing
|
||||
// ============================================================================
|
||||
|
||||
/// List available Gemini models and convert to OpenAI format.
|
||||
///
|
||||
/// Gemini returns `{ models: [{ name: "models/gemini-2.5-flash", displayName, ... }] }`.
|
||||
/// The frontend expects OpenAI format `{ data: [{ id: "models/gemini-2.5-flash", ... }] }`.
|
||||
pub async fn handle_google_ai_models(
|
||||
api_key: &str,
|
||||
base_url: &str,
|
||||
) -> Result<(http::StatusCode, http::HeaderMap, Body)> {
|
||||
#[derive(Deserialize)]
|
||||
struct GeminiModel {
|
||||
name: String,
|
||||
#[serde(rename = "displayName", default)]
|
||||
display_name: String,
|
||||
}
|
||||
|
||||
#[derive(Deserialize)]
|
||||
struct GeminiModelsResponse {
|
||||
#[serde(default)]
|
||||
models: Vec<GeminiModel>,
|
||||
}
|
||||
|
||||
let endpoint = format!("{}/models", base_url.trim_end_matches('/'));
|
||||
let response = HTTP_CLIENT
|
||||
.get(&endpoint)
|
||||
.header("x-goog-api-key", api_key)
|
||||
.send()
|
||||
.await
|
||||
.map_err(|e| Error::internal_err(format!("Failed to fetch Gemini models: {}", e)))?;
|
||||
|
||||
if let Err(e) = response.error_for_status_ref() {
|
||||
let status = e.status().map(|s| s.to_string()).unwrap_or_default();
|
||||
let body = response.text().await.unwrap_or_default();
|
||||
return Err(Error::AIError(format!("{}: {}", status, body)));
|
||||
}
|
||||
|
||||
let gemini_resp: GeminiModelsResponse = response.json().await.map_err(|e| {
|
||||
Error::internal_err(format!("Failed to parse Gemini models response: {}", e))
|
||||
})?;
|
||||
|
||||
let data: Vec<serde_json::Value> = gemini_resp
|
||||
.models
|
||||
.into_iter()
|
||||
.map(|m| {
|
||||
json!({
|
||||
"id": m.name,
|
||||
"object": "model",
|
||||
"display_name": m.display_name,
|
||||
})
|
||||
})
|
||||
.collect();
|
||||
|
||||
let body_bytes = serde_json::to_vec(&json!({ "data": data }))
|
||||
.map_err(|e| Error::internal_err(format!("Failed to serialize models: {}", e)))?;
|
||||
|
||||
let mut headers = http::HeaderMap::new();
|
||||
headers.insert("content-type", "application/json".parse().unwrap());
|
||||
|
||||
Ok((http::StatusCode::OK, headers, Body::from(body_bytes)))
|
||||
}
|
||||
|
||||
// ============================================================================
|
||||
// Non-streaming path
|
||||
// ============================================================================
|
||||
|
||||
async fn handle_non_streaming(
|
||||
model: &str,
|
||||
request_body: String,
|
||||
api_key: &str,
|
||||
base_url: &str,
|
||||
) -> Result<(http::StatusCode, http::HeaderMap, Body)> {
|
||||
let endpoint = format!("{}/models/{}:generateContent", base_url, model);
|
||||
|
||||
let response = HTTP_CLIENT
|
||||
.post(&endpoint)
|
||||
.header("content-type", "application/json")
|
||||
.header("x-goog-api-key", api_key)
|
||||
.body(request_body)
|
||||
.send()
|
||||
.await
|
||||
.map_err(|e| {
|
||||
Error::internal_err(format!("Failed to send request to Gemini API: {}", e))
|
||||
})?;
|
||||
|
||||
if let Err(e) = response.error_for_status_ref() {
|
||||
let status = e.status().map(|s| s.to_string()).unwrap_or_default();
|
||||
let body = response.text().await.unwrap_or_default();
|
||||
return Err(Error::AIError(format!("{}: {}", status, body)));
|
||||
}
|
||||
|
||||
let body = response.bytes().await.map_err(|e| {
|
||||
Error::internal_err(format!("Failed to read Gemini response body: {}", e))
|
||||
})?;
|
||||
|
||||
let parsed = parse_gemini_response(&body)?;
|
||||
let openai_response = gemini_response_to_openai(&parsed, model);
|
||||
|
||||
let body_bytes = serde_json::to_vec(&openai_response)
|
||||
.map_err(|e| Error::internal_err(format!("Failed to serialize response: {}", e)))?;
|
||||
|
||||
let mut headers = http::HeaderMap::new();
|
||||
headers.insert("content-type", "application/json".parse().unwrap());
|
||||
|
||||
Ok((http::StatusCode::OK, headers, Body::from(body_bytes)))
|
||||
}
|
||||
@@ -64,6 +64,7 @@ use crate::scim_oss::has_scim_token;
|
||||
use windmill_common::error::AppError;
|
||||
|
||||
mod ai;
|
||||
mod google;
|
||||
mod apps;
|
||||
pub mod args;
|
||||
mod audit;
|
||||
|
||||
726
backend/windmill-common/src/ai_google.rs
Normal file
726
backend/windmill-common/src/ai_google.rs
Normal file
@@ -0,0 +1,726 @@
|
||||
//! Shared Google AI (Gemini API) types and conversion utilities.
|
||||
//!
|
||||
//! This module provides:
|
||||
//! - Gemini request/response types
|
||||
//! - OpenAI → Gemini message conversion
|
||||
//! - Gemini SSE event parsing
|
||||
//!
|
||||
//! Used by both windmill-api (chat proxy) and windmill-worker (AI agent).
|
||||
|
||||
use serde::{Deserialize, Serialize};
|
||||
|
||||
use crate::ai_types::{ContentPart, ExtraContent, GoogleExtraContent, OpenAIContent, OpenAIMessage, ToolDef, UrlCitation};
|
||||
use crate::error::Error;
|
||||
|
||||
// ============================================================================
|
||||
// Request / Content Types
|
||||
// ============================================================================
|
||||
|
||||
/// Inline data for binary content (images).
|
||||
#[derive(Serialize, Deserialize, Clone, Debug)]
|
||||
pub struct GeminiInlineData {
|
||||
#[serde(rename = "mimeType")]
|
||||
pub mime_type: String,
|
||||
pub data: String,
|
||||
}
|
||||
|
||||
/// A part of content — text, inline data, function call, or function response.
|
||||
#[derive(Serialize, Deserialize, Clone, Debug)]
|
||||
#[serde(untagged)]
|
||||
pub enum GeminiPart {
|
||||
Text {
|
||||
text: String,
|
||||
},
|
||||
InlineData {
|
||||
#[serde(rename = "inlineData")]
|
||||
inline_data: GeminiInlineData,
|
||||
},
|
||||
FunctionCall {
|
||||
#[serde(rename = "functionCall")]
|
||||
function_call: GeminiFunctionCall,
|
||||
/// Thought signature for Gemini 3+ models — required when replaying function calls.
|
||||
#[serde(rename = "thoughtSignature", skip_serializing_if = "Option::is_none")]
|
||||
thought_signature: Option<String>,
|
||||
},
|
||||
FunctionResponse {
|
||||
#[serde(rename = "functionResponse")]
|
||||
function_response: GeminiFunctionResponse,
|
||||
},
|
||||
}
|
||||
|
||||
/// A function call from the model.
|
||||
#[derive(Serialize, Deserialize, Clone, Debug)]
|
||||
pub struct GeminiFunctionCall {
|
||||
pub name: String,
|
||||
pub args: serde_json::Value,
|
||||
}
|
||||
|
||||
/// A function response sent back to the model.
|
||||
#[derive(Serialize, Deserialize, Clone, Debug)]
|
||||
pub struct GeminiFunctionResponse {
|
||||
pub name: String,
|
||||
pub response: serde_json::Value,
|
||||
}
|
||||
|
||||
/// Content message with an optional role and a list of parts.
|
||||
#[derive(Serialize, Clone, Debug)]
|
||||
pub struct GeminiContentMessage {
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub role: Option<String>,
|
||||
pub parts: Vec<GeminiPart>,
|
||||
}
|
||||
|
||||
/// Main request body for `generateContent` / `streamGenerateContent`.
|
||||
#[derive(Serialize)]
|
||||
pub struct GeminiTextRequest {
|
||||
pub contents: Vec<GeminiContentMessage>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub tools: Option<Vec<GeminiTool>>,
|
||||
#[serde(rename = "toolConfig", skip_serializing_if = "Option::is_none")]
|
||||
pub tool_config: Option<GeminiToolConfig>,
|
||||
#[serde(rename = "systemInstruction", skip_serializing_if = "Option::is_none")]
|
||||
pub system_instruction: Option<GeminiContentMessage>,
|
||||
#[serde(rename = "generationConfig", skip_serializing_if = "Option::is_none")]
|
||||
pub generation_config: Option<GeminiGenerationConfig>,
|
||||
}
|
||||
|
||||
/// Tool definition — function declarations and/or Google Search grounding.
|
||||
#[derive(Serialize)]
|
||||
pub struct GeminiTool {
|
||||
#[serde(rename = "functionDeclarations", skip_serializing_if = "Option::is_none")]
|
||||
pub function_declarations: Option<Vec<GeminiFunctionDeclaration>>,
|
||||
#[serde(rename = "googleSearch", skip_serializing_if = "Option::is_none")]
|
||||
pub google_search: Option<serde_json::Value>,
|
||||
}
|
||||
|
||||
/// A single function declaration.
|
||||
///
|
||||
/// `parameters` holds a pre-serialized (and, for the worker, pre-sanitized) JSON Schema.
|
||||
#[derive(Serialize)]
|
||||
pub struct GeminiFunctionDeclaration {
|
||||
pub name: String,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub description: Option<String>,
|
||||
pub parameters: serde_json::Value,
|
||||
}
|
||||
|
||||
/// Tool configuration controlling when and how functions are called.
|
||||
#[derive(Serialize)]
|
||||
pub struct GeminiToolConfig {
|
||||
#[serde(rename = "functionCallingConfig")]
|
||||
pub function_calling_config: GeminiFunctionCallingConfig,
|
||||
}
|
||||
|
||||
/// Function calling mode and optional allow-list.
|
||||
#[derive(Serialize)]
|
||||
pub struct GeminiFunctionCallingConfig {
|
||||
pub mode: String,
|
||||
#[serde(rename = "allowedFunctionNames", skip_serializing_if = "Option::is_none")]
|
||||
pub allowed_function_names: Option<Vec<String>>,
|
||||
}
|
||||
|
||||
/// Generation parameters (temperature, token limits, structured output).
|
||||
#[derive(Serialize)]
|
||||
pub struct GeminiGenerationConfig {
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub temperature: Option<f32>,
|
||||
#[serde(rename = "maxOutputTokens", skip_serializing_if = "Option::is_none")]
|
||||
pub max_output_tokens: Option<u32>,
|
||||
#[serde(rename = "responseMimeType", skip_serializing_if = "Option::is_none")]
|
||||
pub response_mime_type: Option<String>,
|
||||
#[serde(rename = "responseSchema", skip_serializing_if = "Option::is_none")]
|
||||
pub response_schema: Option<serde_json::Value>,
|
||||
}
|
||||
|
||||
// ============================================================================
|
||||
// Image Generation Types
|
||||
// ============================================================================
|
||||
|
||||
/// Request body for Imagen / Gemini image generation.
|
||||
#[derive(Serialize)]
|
||||
pub struct GeminiImageRequest {
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub contents: Option<Vec<GeminiImageContent>>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub instances: Option<Vec<GeminiPredictContent>>,
|
||||
}
|
||||
|
||||
/// Content wrapper used in `generateContent` image requests.
|
||||
#[derive(Serialize)]
|
||||
pub struct GeminiImageContent {
|
||||
pub parts: Vec<GeminiPart>,
|
||||
}
|
||||
|
||||
/// Prompt wrapper for Imagen `predict` endpoint.
|
||||
#[derive(Serialize)]
|
||||
pub struct GeminiPredictContent {
|
||||
pub prompt: String,
|
||||
}
|
||||
|
||||
/// Top-level response from Gemini/Imagen image generation.
|
||||
#[derive(Deserialize)]
|
||||
pub struct GeminiImageResponse {
|
||||
pub candidates: Option<Vec<GeminiImageCandidate>>,
|
||||
pub predictions: Option<Vec<GeminiPredictCandidate>>,
|
||||
}
|
||||
|
||||
#[derive(Deserialize)]
|
||||
pub struct GeminiImageCandidate {
|
||||
pub content: GeminiImageCandidateContent,
|
||||
}
|
||||
|
||||
#[derive(Deserialize)]
|
||||
pub struct GeminiImageCandidateContent {
|
||||
pub parts: Vec<GeminiImageCandidatePart>,
|
||||
}
|
||||
|
||||
#[derive(Deserialize)]
|
||||
pub struct GeminiImageCandidatePart {
|
||||
#[serde(rename = "inlineData")]
|
||||
pub inline_data: Option<GeminiInlineData>,
|
||||
}
|
||||
|
||||
#[derive(Deserialize)]
|
||||
pub struct GeminiPredictCandidate {
|
||||
#[serde(rename = "bytesBase64Encoded")]
|
||||
pub bytes_base64_encoded: String,
|
||||
}
|
||||
|
||||
// ============================================================================
|
||||
// SSE Response Types
|
||||
// ============================================================================
|
||||
|
||||
/// One part inside a streaming candidate — text, function call, or thought signature.
|
||||
#[derive(Deserialize, Debug)]
|
||||
pub struct GeminiSSEPart {
|
||||
#[serde(default)]
|
||||
pub text: Option<String>,
|
||||
#[serde(rename = "functionCall")]
|
||||
pub function_call: Option<GeminiSSEFunctionCall>,
|
||||
/// Thought signature for Gemini 3+ models.
|
||||
#[serde(rename = "thoughtSignature")]
|
||||
pub thought_signature: Option<String>,
|
||||
}
|
||||
|
||||
/// Function call contained in a streaming part.
|
||||
#[derive(Deserialize, Debug)]
|
||||
pub struct GeminiSSEFunctionCall {
|
||||
pub name: String,
|
||||
pub args: serde_json::Value,
|
||||
}
|
||||
|
||||
/// Content block inside a streaming candidate.
|
||||
#[derive(Deserialize, Debug)]
|
||||
pub struct GeminiSSEContent {
|
||||
pub parts: Option<Vec<GeminiSSEPart>>,
|
||||
}
|
||||
|
||||
/// Web source from a Gemini grounding chunk.
|
||||
#[derive(Deserialize, Debug)]
|
||||
pub struct GeminiGroundingChunkWeb {
|
||||
pub uri: String,
|
||||
#[serde(default)]
|
||||
pub title: Option<String>,
|
||||
}
|
||||
|
||||
/// One grounding chunk (search result) from Gemini web search.
|
||||
#[derive(Deserialize, Debug)]
|
||||
pub struct GeminiGroundingChunk {
|
||||
pub web: Option<GeminiGroundingChunkWeb>,
|
||||
}
|
||||
|
||||
/// Grounding metadata attached to a streaming candidate.
|
||||
#[derive(Deserialize, Debug)]
|
||||
pub struct GeminiGroundingMetadata {
|
||||
#[serde(rename = "groundingChunks", default)]
|
||||
pub grounding_chunks: Vec<GeminiGroundingChunk>,
|
||||
#[serde(rename = "webSearchQueries", default)]
|
||||
pub web_search_queries: Vec<String>,
|
||||
}
|
||||
|
||||
/// One candidate inside a streaming Gemini response.
|
||||
#[derive(Deserialize, Debug)]
|
||||
pub struct GeminiSSECandidate {
|
||||
pub content: Option<GeminiSSEContent>,
|
||||
#[serde(rename = "finishReason")]
|
||||
pub finish_reason: Option<String>,
|
||||
#[serde(rename = "groundingMetadata")]
|
||||
pub grounding_metadata: Option<GeminiGroundingMetadata>,
|
||||
}
|
||||
|
||||
/// Token usage from the `usageMetadata` field of a Gemini SSE event.
|
||||
#[derive(Deserialize, Debug, Clone)]
|
||||
pub struct GeminiUsageMetadata {
|
||||
#[serde(rename = "promptTokenCount", default)]
|
||||
pub prompt_token_count: Option<i32>,
|
||||
#[serde(rename = "candidatesTokenCount", default)]
|
||||
pub candidates_token_count: Option<i32>,
|
||||
#[serde(rename = "totalTokenCount", default)]
|
||||
pub total_token_count: Option<i32>,
|
||||
}
|
||||
|
||||
/// Top-level structure of one Gemini SSE event.
|
||||
#[derive(Deserialize, Debug)]
|
||||
pub struct GeminiSSEEvent {
|
||||
pub candidates: Option<Vec<GeminiSSECandidate>>,
|
||||
#[serde(rename = "usageMetadata")]
|
||||
pub usage_metadata: Option<GeminiUsageMetadata>,
|
||||
}
|
||||
|
||||
// ============================================================================
|
||||
// Parsed Event Result
|
||||
// ============================================================================
|
||||
|
||||
/// A single function call extracted from a Gemini SSE event.
|
||||
#[derive(Debug)]
|
||||
pub struct GeminiToolCallEvent {
|
||||
pub name: String,
|
||||
pub args: serde_json::Value,
|
||||
pub thought_signature: Option<String>,
|
||||
}
|
||||
|
||||
impl GeminiToolCallEvent {
|
||||
/// Convert the thought signature (if present) into an [`ExtraContent`].
|
||||
pub fn to_extra_content(&self) -> Option<ExtraContent> {
|
||||
self.thought_signature.as_ref().map(|sig| ExtraContent {
|
||||
google: Some(GoogleExtraContent { thought_signature: Some(sig.clone()) }),
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
/// Structured result of parsing a Gemini response (streaming SSE event or non-streaming body).
|
||||
#[derive(Debug, Default)]
|
||||
pub struct GeminiParsedEvent {
|
||||
pub text: Option<String>,
|
||||
pub tool_calls: Vec<GeminiToolCallEvent>,
|
||||
pub annotations: Vec<UrlCitation>,
|
||||
pub used_websearch: bool,
|
||||
pub usage: Option<GeminiUsageMetadata>,
|
||||
pub finish_reason: Option<String>,
|
||||
}
|
||||
|
||||
// ============================================================================
|
||||
// Helper Functions
|
||||
// ============================================================================
|
||||
|
||||
/// Parse a data URL into `(mime_type, base64_data)`.
|
||||
///
|
||||
/// Expected format: `data:<mime_type>;base64,<data>`.
|
||||
pub fn parse_data_url(url: &str) -> Option<(String, String)> {
|
||||
let rest = url.strip_prefix("data:")?;
|
||||
let (header, data) = rest.split_once(',')?;
|
||||
let media_type = header.strip_suffix(";base64")?;
|
||||
Some((media_type.to_string(), data.to_string()))
|
||||
}
|
||||
|
||||
/// Find the function name associated with a `tool_call_id` by scanning prior messages.
|
||||
pub fn find_gemini_function_name(messages: &[OpenAIMessage], tool_call_id: &str) -> String {
|
||||
messages
|
||||
.iter()
|
||||
.filter_map(|msg| msg.tool_calls.as_ref())
|
||||
.flatten()
|
||||
.find(|tc| tc.id == tool_call_id)
|
||||
.map(|tc| tc.function.name.clone())
|
||||
.unwrap_or_else(|| "unknown_function".to_string())
|
||||
}
|
||||
|
||||
/// Convert an [`OpenAIContent`] value to a list of [`GeminiPart`]s.
|
||||
///
|
||||
/// Handles text and `image_url` (data URLs). `S3Object` variants are skipped here;
|
||||
/// the worker handles them by downloading and injecting inline data beforehand.
|
||||
pub fn convert_content_to_gemini_parts(content: &OpenAIContent) -> Vec<GeminiPart> {
|
||||
match content {
|
||||
OpenAIContent::Text(text) if !text.is_empty() => {
|
||||
vec![GeminiPart::Text { text: text.clone() }]
|
||||
}
|
||||
OpenAIContent::Text(_) => vec![],
|
||||
OpenAIContent::Parts(parts) => parts
|
||||
.iter()
|
||||
.filter_map(|part| match part {
|
||||
ContentPart::Text { text } if !text.is_empty() => {
|
||||
Some(GeminiPart::Text { text: text.clone() })
|
||||
}
|
||||
ContentPart::ImageUrl { image_url } => {
|
||||
parse_data_url(&image_url.url).map(|(mime_type, data)| {
|
||||
GeminiPart::InlineData {
|
||||
inline_data: GeminiInlineData { mime_type, data },
|
||||
}
|
||||
})
|
||||
}
|
||||
// S3Objects are handled by the worker
|
||||
_ => None,
|
||||
})
|
||||
.collect(),
|
||||
}
|
||||
}
|
||||
|
||||
/// Convert OpenAI-format messages to Gemini `contents` and an optional `systemInstruction`.
|
||||
///
|
||||
/// Returns `(contents, system_instruction)`.
|
||||
///
|
||||
/// `S3Object` images in content parts are skipped (the worker pre-converts them).
|
||||
/// Tool call history is preserved correctly for multi-turn agent conversations.
|
||||
pub fn openai_messages_to_gemini(
|
||||
messages: &[OpenAIMessage],
|
||||
) -> (Vec<GeminiContentMessage>, Option<GeminiContentMessage>) {
|
||||
let mut contents: Vec<GeminiContentMessage> = Vec::new();
|
||||
let mut system_instruction: Option<GeminiContentMessage> = None;
|
||||
|
||||
for msg in messages {
|
||||
match msg.role.as_str() {
|
||||
"system" => {
|
||||
if let Some(content) = &msg.content {
|
||||
let parts = convert_content_to_gemini_parts(content);
|
||||
if !parts.is_empty() {
|
||||
system_instruction =
|
||||
Some(GeminiContentMessage { role: None, parts });
|
||||
}
|
||||
}
|
||||
}
|
||||
"tool" => {
|
||||
if let (Some(tool_call_id), Some(content)) =
|
||||
(&msg.tool_call_id, &msg.content)
|
||||
{
|
||||
let func_name = find_gemini_function_name(messages, tool_call_id);
|
||||
let response_text = match content {
|
||||
OpenAIContent::Text(text) => text.clone(),
|
||||
OpenAIContent::Parts(parts) => parts
|
||||
.iter()
|
||||
.filter_map(|p| {
|
||||
if let ContentPart::Text { text } = p {
|
||||
Some(text.as_str())
|
||||
} else {
|
||||
None
|
||||
}
|
||||
})
|
||||
.collect::<Vec<_>>()
|
||||
.join(" "),
|
||||
};
|
||||
contents.push(GeminiContentMessage {
|
||||
role: Some("user".to_string()),
|
||||
parts: vec![GeminiPart::FunctionResponse {
|
||||
function_response: GeminiFunctionResponse {
|
||||
name: func_name,
|
||||
response: serde_json::json!({ "result": response_text }),
|
||||
},
|
||||
}],
|
||||
});
|
||||
}
|
||||
}
|
||||
role => {
|
||||
let gemini_role = if role == "assistant" { "model" } else { "user" };
|
||||
let mut parts: Vec<GeminiPart> = Vec::new();
|
||||
|
||||
if let Some(content) = &msg.content {
|
||||
parts.extend(convert_content_to_gemini_parts(content));
|
||||
}
|
||||
|
||||
if let Some(tool_calls) = &msg.tool_calls {
|
||||
for tc in tool_calls {
|
||||
let args: serde_json::Value =
|
||||
serde_json::from_str(&tc.function.arguments).unwrap_or_default();
|
||||
let thought_signature = tc
|
||||
.extra_content
|
||||
.as_ref()
|
||||
.and_then(|ec| ec.google.as_ref())
|
||||
.and_then(|g| g.thought_signature.clone());
|
||||
parts.push(GeminiPart::FunctionCall {
|
||||
function_call: GeminiFunctionCall {
|
||||
name: tc.function.name.clone(),
|
||||
args,
|
||||
},
|
||||
thought_signature,
|
||||
});
|
||||
}
|
||||
}
|
||||
|
||||
if !parts.is_empty() {
|
||||
contents.push(GeminiContentMessage {
|
||||
role: Some(gemini_role.to_string()),
|
||||
parts,
|
||||
});
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
(contents, system_instruction)
|
||||
}
|
||||
|
||||
/// Convert OpenAI tool definitions to Gemini format.
|
||||
///
|
||||
/// `tool_params` must be pre-serialized (and, for the worker, pre-sanitized for Google)
|
||||
/// JSON schema values, one per entry in `tools` in the same order.
|
||||
pub fn openai_tools_to_gemini(
|
||||
tools: &[ToolDef],
|
||||
tool_params: &[serde_json::Value],
|
||||
has_websearch: bool,
|
||||
) -> Option<Vec<GeminiTool>> {
|
||||
let mut gemini_tools: Vec<GeminiTool> = Vec::new();
|
||||
|
||||
let declarations: Vec<GeminiFunctionDeclaration> = tools
|
||||
.iter()
|
||||
.zip(tool_params.iter())
|
||||
.map(|(t, params)| GeminiFunctionDeclaration {
|
||||
name: t.function.name.clone(),
|
||||
description: t.function.description.clone(),
|
||||
parameters: params.clone(),
|
||||
})
|
||||
.collect();
|
||||
|
||||
if !declarations.is_empty() {
|
||||
gemini_tools.push(GeminiTool {
|
||||
function_declarations: Some(declarations),
|
||||
google_search: None,
|
||||
});
|
||||
}
|
||||
|
||||
if has_websearch {
|
||||
gemini_tools.push(GeminiTool {
|
||||
function_declarations: None,
|
||||
google_search: Some(serde_json::json!({})),
|
||||
});
|
||||
}
|
||||
|
||||
if gemini_tools.is_empty() {
|
||||
None
|
||||
} else {
|
||||
Some(gemini_tools)
|
||||
}
|
||||
}
|
||||
|
||||
/// Parse one Gemini SSE data line into a [`GeminiParsedEvent`].
|
||||
///
|
||||
/// Returns `Ok(None)` for empty data or unrecognised payloads (e.g. `"[DONE]"`).
|
||||
/// Logs a warning and returns `Ok(None)` on JSON parse errors rather than propagating.
|
||||
pub fn parse_gemini_sse_event(data: &str) -> Result<Option<GeminiParsedEvent>, Error> {
|
||||
if data.is_empty() || data == "[DONE]" {
|
||||
return Ok(None);
|
||||
}
|
||||
|
||||
let event: GeminiSSEEvent = match serde_json::from_str(data) {
|
||||
Ok(e) => e,
|
||||
Err(e) => {
|
||||
tracing::error!("Failed to parse Gemini SSE event {}: {}", data, e);
|
||||
return Ok(None);
|
||||
}
|
||||
};
|
||||
|
||||
let mut parsed = GeminiParsedEvent { usage: event.usage_metadata, ..Default::default() };
|
||||
|
||||
let Some(candidates) = event.candidates else {
|
||||
return Ok(Some(parsed));
|
||||
};
|
||||
|
||||
extract_candidates_into(&candidates, &mut parsed);
|
||||
|
||||
Ok(Some(parsed))
|
||||
}
|
||||
|
||||
/// Parse a non-streaming Gemini `generateContent` response body.
|
||||
pub fn parse_gemini_response(data: &[u8]) -> Result<GeminiParsedEvent, Error> {
|
||||
let event: GeminiSSEEvent = serde_json::from_slice(data)
|
||||
.map_err(|e| Error::internal_err(format!("Failed to parse Gemini response: {}", e)))?;
|
||||
|
||||
let mut parsed = GeminiParsedEvent { usage: event.usage_metadata, ..Default::default() };
|
||||
|
||||
if let Some(candidates) = event.candidates {
|
||||
extract_candidates_into(&candidates, &mut parsed);
|
||||
}
|
||||
|
||||
Ok(parsed)
|
||||
}
|
||||
|
||||
// ============================================================================
|
||||
// Gemini → OpenAI Format Conversion
|
||||
// ============================================================================
|
||||
|
||||
/// Convert a `GeminiParsedEvent` from a non-streaming response to an OpenAI chat completion JSON.
|
||||
pub fn gemini_response_to_openai(parsed: &GeminiParsedEvent, model: &str) -> serde_json::Value {
|
||||
let content = parsed.text.as_deref().unwrap_or_default();
|
||||
|
||||
let tool_calls: Vec<serde_json::Value> = parsed
|
||||
.tool_calls
|
||||
.iter()
|
||||
.enumerate()
|
||||
.map(|(i, tc)| {
|
||||
serde_json::json!({
|
||||
"index": i,
|
||||
"id": format!("call_{}", uuid::Uuid::new_v4().simple()),
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": tc.name,
|
||||
"arguments": serde_json::to_string(&tc.args).unwrap_or_default()
|
||||
}
|
||||
})
|
||||
})
|
||||
.collect();
|
||||
|
||||
let finish_reason = parsed
|
||||
.finish_reason
|
||||
.as_deref()
|
||||
.map(|r| r.to_lowercase())
|
||||
.unwrap_or_else(|| "stop".to_string());
|
||||
|
||||
let usage = parsed.usage.as_ref().map(|u| {
|
||||
serde_json::json!({
|
||||
"prompt_tokens": u.prompt_token_count.unwrap_or(0),
|
||||
"completion_tokens": u.candidates_token_count.unwrap_or(0),
|
||||
"total_tokens": u.total_token_count.unwrap_or(0),
|
||||
})
|
||||
});
|
||||
|
||||
let mut message = serde_json::json!({
|
||||
"role": "assistant",
|
||||
"content": content,
|
||||
});
|
||||
if !tool_calls.is_empty() {
|
||||
message["tool_calls"] = serde_json::json!(tool_calls);
|
||||
}
|
||||
|
||||
serde_json::json!({
|
||||
"id": format!("chatcmpl-{}", uuid::Uuid::new_v4().simple()),
|
||||
"object": "chat.completion",
|
||||
"model": model,
|
||||
"choices": [{
|
||||
"index": 0,
|
||||
"message": message,
|
||||
"finish_reason": finish_reason,
|
||||
}],
|
||||
"usage": usage,
|
||||
})
|
||||
}
|
||||
|
||||
/// Convert a `GeminiParsedEvent` from a streaming SSE event into OpenAI-format SSE lines.
|
||||
///
|
||||
/// Returns the serialized `"data: {...}\n\n"` lines ready to be written to the response stream.
|
||||
/// `tool_call_index` is mutated to track the running index across multiple SSE events.
|
||||
pub fn gemini_event_to_openai_sse_chunks(
|
||||
parsed: &GeminiParsedEvent,
|
||||
id: &str,
|
||||
model: &str,
|
||||
tool_call_index: &mut usize,
|
||||
) -> Vec<String> {
|
||||
let mut chunks = Vec::new();
|
||||
|
||||
if let Some(text) = &parsed.text {
|
||||
let chunk = serde_json::json!({
|
||||
"id": id,
|
||||
"object": "chat.completion.chunk",
|
||||
"model": model,
|
||||
"choices": [{
|
||||
"index": 0,
|
||||
"delta": { "content": text },
|
||||
"finish_reason": null,
|
||||
}]
|
||||
});
|
||||
chunks.push(format!("data: {}\n\n", chunk));
|
||||
}
|
||||
|
||||
for tc in &parsed.tool_calls {
|
||||
let args_str = serde_json::to_string(&tc.args).unwrap_or_default();
|
||||
let call_id = format!("call_{}", uuid::Uuid::new_v4().simple());
|
||||
let chunk = serde_json::json!({
|
||||
"id": id,
|
||||
"object": "chat.completion.chunk",
|
||||
"model": model,
|
||||
"choices": [{
|
||||
"index": 0,
|
||||
"delta": {
|
||||
"tool_calls": [{
|
||||
"index": *tool_call_index,
|
||||
"id": call_id,
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": tc.name,
|
||||
"arguments": args_str,
|
||||
}
|
||||
}]
|
||||
},
|
||||
"finish_reason": null,
|
||||
}]
|
||||
});
|
||||
chunks.push(format!("data: {}\n\n", chunk));
|
||||
*tool_call_index += 1;
|
||||
}
|
||||
|
||||
chunks
|
||||
}
|
||||
|
||||
/// Recursively remove JSON Schema fields unsupported by the Gemini API.
|
||||
pub fn sanitize_schema_for_google(value: &mut serde_json::Value) {
|
||||
const UNSUPPORTED: &[&str] = &[
|
||||
"additionalProperties",
|
||||
"strict",
|
||||
"$schema",
|
||||
"default",
|
||||
"exclusiveMinimum",
|
||||
"exclusiveMaximum",
|
||||
"const",
|
||||
"multipleOf",
|
||||
];
|
||||
|
||||
if let Some(obj) = value.as_object_mut() {
|
||||
for field in UNSUPPORTED {
|
||||
obj.remove(*field);
|
||||
}
|
||||
for v in obj.values_mut() {
|
||||
sanitize_schema_for_google(v);
|
||||
}
|
||||
} else if let Some(arr) = value.as_array_mut() {
|
||||
for v in arr.iter_mut() {
|
||||
sanitize_schema_for_google(v);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// ============================================================================
|
||||
// Internal Helpers
|
||||
// ============================================================================
|
||||
|
||||
fn extract_candidates_into(candidates: &[GeminiSSECandidate], parsed: &mut GeminiParsedEvent) {
|
||||
for candidate in candidates {
|
||||
if let Some(content) = &candidate.content {
|
||||
if let Some(parts) = &content.parts {
|
||||
for part in parts {
|
||||
if let Some(text) = &part.text {
|
||||
if !text.is_empty() {
|
||||
match parsed.text.as_mut() {
|
||||
Some(existing) => existing.push_str(text),
|
||||
None => parsed.text = Some(text.clone()),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
if let Some(function_call) = &part.function_call {
|
||||
parsed.tool_calls.push(GeminiToolCallEvent {
|
||||
name: function_call.name.clone(),
|
||||
args: function_call.args.clone(),
|
||||
thought_signature: part.thought_signature.clone(),
|
||||
});
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
if candidate.finish_reason.is_some() {
|
||||
parsed.finish_reason = candidate.finish_reason.clone();
|
||||
}
|
||||
|
||||
if let Some(grounding) = &candidate.grounding_metadata {
|
||||
if !grounding.web_search_queries.is_empty() || !grounding.grounding_chunks.is_empty() {
|
||||
parsed.used_websearch = true;
|
||||
}
|
||||
for chunk in &grounding.grounding_chunks {
|
||||
if let Some(web) = &chunk.web {
|
||||
parsed.annotations.push(UrlCitation {
|
||||
start_index: 0,
|
||||
end_index: 0,
|
||||
url: web.uri.clone(),
|
||||
title: web.title.clone(),
|
||||
});
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -29,6 +29,7 @@ use sqlx::{Acquire, Postgres};
|
||||
pub mod agent_workers;
|
||||
#[cfg(feature = "bedrock")]
|
||||
pub mod ai_bedrock;
|
||||
pub mod ai_google;
|
||||
pub mod ai_providers;
|
||||
pub mod ai_types;
|
||||
pub mod apps;
|
||||
|
||||
@@ -1,14 +1,16 @@
|
||||
use async_trait::async_trait;
|
||||
use serde::{Deserialize, Serialize};
|
||||
use serde_json::value::RawValue;
|
||||
use windmill_common::{ai_providers::AIProvider, client::AuthedClient, error::Error};
|
||||
use windmill_common::{
|
||||
ai_google::parse_data_url, ai_providers::AIProvider, client::AuthedClient, error::Error,
|
||||
};
|
||||
|
||||
use crate::ai::{
|
||||
image_handler::prepare_messages_for_api,
|
||||
query_builder::{BuildRequestArgs, ParsedResponse, QueryBuilder, StreamEventProcessor},
|
||||
sse::{AnthropicSSEParser, SSEParser},
|
||||
types::*,
|
||||
utils::{extract_text_content, parse_data_url, should_use_structured_output_tool},
|
||||
utils::{extract_text_content, should_use_structured_output_tool},
|
||||
};
|
||||
|
||||
/// Anthropic API version for standard API
|
||||
|
||||
@@ -1,215 +1,21 @@
|
||||
use async_trait::async_trait;
|
||||
use serde::{Deserialize, Serialize};
|
||||
use windmill_common::{client::AuthedClient, error::Error};
|
||||
use windmill_common::{
|
||||
ai_google::{
|
||||
openai_messages_to_gemini, openai_tools_to_gemini, GeminiGenerationConfig,
|
||||
GeminiImageContent, GeminiImageRequest, GeminiImageResponse, GeminiInlineData, GeminiPart,
|
||||
GeminiPredictContent, GeminiTextRequest, GeminiTool,
|
||||
},
|
||||
client::AuthedClient,
|
||||
error::Error,
|
||||
};
|
||||
|
||||
use crate::ai::{
|
||||
image_handler::download_and_encode_s3_image,
|
||||
image_handler::{download_and_encode_s3_image, prepare_messages_for_api},
|
||||
query_builder::{BuildRequestArgs, ParsedResponse, QueryBuilder, StreamEventProcessor},
|
||||
sse::{GeminiSSEParser, SSEParser},
|
||||
types::*,
|
||||
utils::parse_data_url,
|
||||
};
|
||||
|
||||
// ============================================================================
|
||||
// Gemini API Types - Shared between text and image
|
||||
// ============================================================================
|
||||
|
||||
/// Inline data for binary content (images)
|
||||
#[derive(Serialize, Deserialize, Clone, Debug)]
|
||||
pub struct GeminiInlineData {
|
||||
#[serde(rename = "mimeType")]
|
||||
pub mime_type: String,
|
||||
pub data: String,
|
||||
}
|
||||
|
||||
/// A part of content - can be text, inline data, function call, or function response
|
||||
#[derive(Serialize, Deserialize, Clone, Debug)]
|
||||
#[serde(untagged)]
|
||||
pub enum GeminiPart {
|
||||
Text {
|
||||
text: String,
|
||||
},
|
||||
InlineData {
|
||||
#[serde(rename = "inlineData")]
|
||||
inline_data: GeminiInlineData,
|
||||
},
|
||||
FunctionCall {
|
||||
#[serde(rename = "functionCall")]
|
||||
function_call: GeminiFunctionCall,
|
||||
/// Thought signature for Gemini 3+ models - required for function calling
|
||||
#[serde(rename = "thoughtSignature", skip_serializing_if = "Option::is_none")]
|
||||
thought_signature: Option<String>,
|
||||
},
|
||||
FunctionResponse {
|
||||
#[serde(rename = "functionResponse")]
|
||||
function_response: GeminiFunctionResponse,
|
||||
},
|
||||
}
|
||||
|
||||
/// A function call from the model
|
||||
#[derive(Serialize, Deserialize, Clone, Debug)]
|
||||
pub struct GeminiFunctionCall {
|
||||
pub name: String,
|
||||
pub args: serde_json::Value,
|
||||
}
|
||||
|
||||
/// A function response to send back to the model
|
||||
#[derive(Serialize, Deserialize, Clone, Debug)]
|
||||
pub struct GeminiFunctionResponse {
|
||||
pub name: String,
|
||||
pub response: serde_json::Value,
|
||||
}
|
||||
|
||||
// ============================================================================
|
||||
// Gemini Text API Request Types
|
||||
// ============================================================================
|
||||
|
||||
/// Main request structure for Gemini generateContent
|
||||
#[derive(Serialize)]
|
||||
pub struct GeminiTextRequest {
|
||||
pub contents: Vec<GeminiContentMessage>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub tools: Option<Vec<GeminiTool>>,
|
||||
#[serde(rename = "toolConfig", skip_serializing_if = "Option::is_none")]
|
||||
pub tool_config: Option<GeminiToolConfig>,
|
||||
#[serde(rename = "systemInstruction", skip_serializing_if = "Option::is_none")]
|
||||
pub system_instruction: Option<GeminiContentMessage>,
|
||||
#[serde(rename = "generationConfig", skip_serializing_if = "Option::is_none")]
|
||||
pub generation_config: Option<GeminiGenerationConfig>,
|
||||
}
|
||||
|
||||
/// Content message with role and parts
|
||||
#[derive(Serialize)]
|
||||
pub struct GeminiContentMessage {
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub role: Option<String>,
|
||||
pub parts: Vec<GeminiPart>,
|
||||
}
|
||||
|
||||
/// Tool definition - either function declarations or Google Search
|
||||
#[derive(Serialize)]
|
||||
pub struct GeminiTool {
|
||||
#[serde(
|
||||
rename = "functionDeclarations",
|
||||
skip_serializing_if = "Option::is_none"
|
||||
)]
|
||||
pub function_declarations: Option<Vec<GeminiFunctionDeclaration>>,
|
||||
#[serde(rename = "googleSearch", skip_serializing_if = "Option::is_none")]
|
||||
pub google_search: Option<serde_json::Value>,
|
||||
}
|
||||
|
||||
/// Function declaration for tool use
|
||||
#[derive(Serialize)]
|
||||
pub struct GeminiFunctionDeclaration {
|
||||
pub name: String,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub description: Option<String>,
|
||||
pub parameters: OpenAPISchema,
|
||||
}
|
||||
|
||||
/// Tool configuration for controlling function calling behavior
|
||||
#[derive(Serialize)]
|
||||
pub struct GeminiToolConfig {
|
||||
#[serde(rename = "functionCallingConfig")]
|
||||
pub function_calling_config: GeminiFunctionCallingConfig,
|
||||
}
|
||||
|
||||
/// Function calling configuration
|
||||
#[derive(Serialize)]
|
||||
pub struct GeminiFunctionCallingConfig {
|
||||
pub mode: String,
|
||||
#[serde(
|
||||
rename = "allowedFunctionNames",
|
||||
skip_serializing_if = "Option::is_none"
|
||||
)]
|
||||
pub allowed_function_names: Option<Vec<String>>,
|
||||
}
|
||||
|
||||
/// Generation configuration for output format
|
||||
#[derive(Serialize)]
|
||||
pub struct GeminiGenerationConfig {
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub temperature: Option<f32>,
|
||||
#[serde(rename = "maxOutputTokens", skip_serializing_if = "Option::is_none")]
|
||||
pub max_output_tokens: Option<u32>,
|
||||
#[serde(rename = "responseMimeType", skip_serializing_if = "Option::is_none")]
|
||||
pub response_mime_type: Option<String>,
|
||||
#[serde(rename = "responseSchema", skip_serializing_if = "Option::is_none")]
|
||||
pub response_schema: Option<serde_json::Value>,
|
||||
}
|
||||
|
||||
// ============================================================================
|
||||
// Gemini API Response Types
|
||||
// ============================================================================
|
||||
|
||||
/// Grounding metadata from Google Search
|
||||
#[derive(Deserialize)]
|
||||
#[allow(dead_code)]
|
||||
pub struct GeminiGroundingMetadata {
|
||||
#[serde(rename = "webSearchQueries")]
|
||||
pub web_search_queries: Option<Vec<String>>,
|
||||
#[serde(rename = "groundingChunks")]
|
||||
pub grounding_chunks: Option<Vec<serde_json::Value>>,
|
||||
}
|
||||
|
||||
// ============================================================================
|
||||
// Gemini Image API Types (for Imagen models)
|
||||
// ============================================================================
|
||||
|
||||
/// Request for image generation (Imagen models)
|
||||
#[derive(Serialize)]
|
||||
pub struct GeminiImageRequest {
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub contents: Option<Vec<GeminiImageContent>>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub instances: Option<Vec<GeminiPredictContent>>,
|
||||
}
|
||||
|
||||
/// Content for image generation
|
||||
#[derive(Serialize)]
|
||||
pub struct GeminiImageContent {
|
||||
pub parts: Vec<GeminiPart>,
|
||||
}
|
||||
|
||||
/// Content for Imagen predict endpoint
|
||||
#[derive(Serialize)]
|
||||
pub struct GeminiPredictContent {
|
||||
pub prompt: String,
|
||||
}
|
||||
|
||||
/// Response for image generation
|
||||
#[derive(Deserialize)]
|
||||
pub struct GeminiImageResponse {
|
||||
pub candidates: Option<Vec<GeminiImageCandidate>>,
|
||||
pub predictions: Option<Vec<GeminiPredictCandidate>>,
|
||||
}
|
||||
|
||||
/// Image candidate from generateContent
|
||||
#[derive(Deserialize)]
|
||||
pub struct GeminiImageCandidate {
|
||||
pub content: GeminiImageCandidateContent,
|
||||
}
|
||||
|
||||
/// Content in image candidate
|
||||
#[derive(Deserialize)]
|
||||
pub struct GeminiImageCandidateContent {
|
||||
pub parts: Vec<GeminiImageCandidatePart>,
|
||||
}
|
||||
|
||||
/// Part of image candidate
|
||||
#[derive(Deserialize)]
|
||||
pub struct GeminiImageCandidatePart {
|
||||
#[serde(rename = "inlineData", skip_serializing_if = "Option::is_none")]
|
||||
pub inline_data: Option<GeminiInlineData>,
|
||||
}
|
||||
|
||||
/// Prediction candidate from Imagen
|
||||
#[derive(Deserialize)]
|
||||
pub struct GeminiPredictCandidate {
|
||||
#[serde(rename = "bytesBase64Encoded")]
|
||||
pub bytes_base64_encoded: String,
|
||||
}
|
||||
|
||||
// ============================================================================
|
||||
// Query Builder Implementation
|
||||
// ============================================================================
|
||||
@@ -221,34 +27,24 @@ impl GoogleAIQueryBuilder {
|
||||
Self
|
||||
}
|
||||
|
||||
/// Build a text request using the native Gemini API format
|
||||
async fn build_text_request(
|
||||
&self,
|
||||
args: &BuildRequestArgs<'_>,
|
||||
client: &AuthedClient,
|
||||
workspace_id: &str,
|
||||
) -> Result<String, Error> {
|
||||
// Convert messages to Gemini format
|
||||
let contents = self
|
||||
.convert_messages_to_gemini(args.messages, client, workspace_id)
|
||||
.await?;
|
||||
let prepared_messages =
|
||||
prepare_messages_for_api(args.messages, client, workspace_id).await?;
|
||||
let (contents, system_instruction) = openai_messages_to_gemini(&prepared_messages);
|
||||
|
||||
// Build tools array
|
||||
let tools = self.convert_tools_to_gemini(args.tools, args.has_websearch);
|
||||
|
||||
// Build generation config
|
||||
let generation_config = self.build_generation_config(args);
|
||||
|
||||
// Build system instruction from system_prompt
|
||||
let system_instruction = args.system_prompt.map(|s| GeminiContentMessage {
|
||||
role: None,
|
||||
parts: vec![GeminiPart::Text { text: s.to_string() }],
|
||||
});
|
||||
|
||||
let request = GeminiTextRequest {
|
||||
contents,
|
||||
tools,
|
||||
tool_config: None, // Use AUTO mode by default
|
||||
tool_config: None,
|
||||
system_instruction,
|
||||
generation_config,
|
||||
};
|
||||
@@ -257,7 +53,6 @@ impl GoogleAIQueryBuilder {
|
||||
.map_err(|e| Error::internal_err(format!("Failed to serialize request: {}", e)))
|
||||
}
|
||||
|
||||
/// Build an image generation request
|
||||
async fn build_image_request(
|
||||
&self,
|
||||
args: &BuildRequestArgs<'_>,
|
||||
@@ -267,7 +62,6 @@ impl GoogleAIQueryBuilder {
|
||||
let is_imagen = args.model.contains("imagen");
|
||||
|
||||
let request = if is_imagen {
|
||||
// For Imagen models, use simple prompt format
|
||||
GeminiImageRequest {
|
||||
instances: Some(vec![GeminiPredictContent {
|
||||
prompt: args.user_message.trim().to_string(),
|
||||
@@ -275,7 +69,6 @@ impl GoogleAIQueryBuilder {
|
||||
contents: None,
|
||||
}
|
||||
} else {
|
||||
// For Gemini models with image generation, build parts
|
||||
let mut parts = vec![GeminiPart::Text { text: args.user_message.trim().to_string() }];
|
||||
|
||||
if let Some(system_prompt) = args.system_prompt {
|
||||
@@ -285,7 +78,6 @@ impl GoogleAIQueryBuilder {
|
||||
);
|
||||
}
|
||||
|
||||
// Add input images if provided
|
||||
if let Some(images) = args.images {
|
||||
for image in images.iter() {
|
||||
if !image.s3.is_empty() {
|
||||
@@ -308,218 +100,39 @@ impl GoogleAIQueryBuilder {
|
||||
.map_err(|e| Error::internal_err(format!("Failed to serialize request: {}", e)))
|
||||
}
|
||||
|
||||
/// Convert OpenAI-format messages to Gemini format
|
||||
async fn convert_messages_to_gemini(
|
||||
&self,
|
||||
messages: &[OpenAIMessage],
|
||||
client: &AuthedClient,
|
||||
workspace_id: &str,
|
||||
) -> Result<Vec<GeminiContentMessage>, Error> {
|
||||
let mut gemini_messages = Vec::new();
|
||||
|
||||
for msg in messages {
|
||||
match msg.role.as_str() {
|
||||
"system" => {
|
||||
// Skip - handled via args.system_prompt in build_text_request
|
||||
}
|
||||
"tool" => {
|
||||
// Handle tool responses
|
||||
if let (Some(tool_call_id), Some(content)) = (&msg.tool_call_id, &msg.content) {
|
||||
let func_name = self.find_function_name_by_id(messages, tool_call_id);
|
||||
let response_text = match content {
|
||||
OpenAIContent::Text(text) => text.clone(),
|
||||
OpenAIContent::Parts(parts) => parts
|
||||
.iter()
|
||||
.filter_map(|p| match p {
|
||||
ContentPart::Text { text } => Some(text.clone()),
|
||||
_ => None,
|
||||
})
|
||||
.collect::<Vec<_>>()
|
||||
.join(" "),
|
||||
};
|
||||
|
||||
gemini_messages.push(GeminiContentMessage {
|
||||
role: Some("user".to_string()),
|
||||
parts: vec![GeminiPart::FunctionResponse {
|
||||
function_response: GeminiFunctionResponse {
|
||||
name: func_name,
|
||||
response: serde_json::json!({ "result": response_text }),
|
||||
},
|
||||
}],
|
||||
});
|
||||
}
|
||||
}
|
||||
_ => {
|
||||
// Handle user/assistant messages
|
||||
let role = match msg.role.as_str() {
|
||||
"assistant" => "model",
|
||||
_ => "user",
|
||||
};
|
||||
|
||||
let mut parts = Vec::new();
|
||||
|
||||
// Handle regular content
|
||||
if let Some(content) = &msg.content {
|
||||
let content_parts = self
|
||||
.convert_content_to_parts(&Some(content.clone()), client, workspace_id)
|
||||
.await?;
|
||||
parts.extend(content_parts);
|
||||
}
|
||||
|
||||
// Handle tool calls from assistant
|
||||
if let Some(tool_calls) = &msg.tool_calls {
|
||||
for tc in tool_calls {
|
||||
let args: serde_json::Value =
|
||||
serde_json::from_str(&tc.function.arguments).unwrap_or_default();
|
||||
// Extract thought_signature from extra_content if present
|
||||
let thought_signature = tc
|
||||
.extra_content
|
||||
.as_ref()
|
||||
.and_then(|ec| ec.google.as_ref())
|
||||
.and_then(|g| g.thought_signature.clone());
|
||||
parts.push(GeminiPart::FunctionCall {
|
||||
function_call: GeminiFunctionCall {
|
||||
name: tc.function.name.clone(),
|
||||
args,
|
||||
},
|
||||
thought_signature,
|
||||
});
|
||||
}
|
||||
}
|
||||
|
||||
if !parts.is_empty() {
|
||||
gemini_messages
|
||||
.push(GeminiContentMessage { role: Some(role.to_string()), parts });
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
Ok(gemini_messages)
|
||||
}
|
||||
|
||||
/// Convert OpenAI content to Gemini parts
|
||||
async fn convert_content_to_parts(
|
||||
&self,
|
||||
content: &Option<OpenAIContent>,
|
||||
client: &AuthedClient,
|
||||
workspace_id: &str,
|
||||
) -> Result<Vec<GeminiPart>, Error> {
|
||||
let mut parts = Vec::new();
|
||||
|
||||
if let Some(content) = content {
|
||||
match content {
|
||||
OpenAIContent::Text(text) => {
|
||||
if !text.is_empty() {
|
||||
parts.push(GeminiPart::Text { text: text.clone() });
|
||||
}
|
||||
}
|
||||
OpenAIContent::Parts(content_parts) => {
|
||||
for part in content_parts {
|
||||
match part {
|
||||
ContentPart::Text { text } => {
|
||||
if !text.is_empty() {
|
||||
parts.push(GeminiPart::Text { text: text.clone() });
|
||||
}
|
||||
}
|
||||
ContentPart::ImageUrl { image_url } => {
|
||||
// Parse data URL format: data:mime_type;base64,data
|
||||
if let Some((mime_type, data)) = parse_data_url(&image_url.url) {
|
||||
parts.push(GeminiPart::InlineData {
|
||||
inline_data: GeminiInlineData { mime_type, data },
|
||||
});
|
||||
}
|
||||
}
|
||||
ContentPart::S3Object { s3_object } => {
|
||||
if !s3_object.s3.is_empty() {
|
||||
let (mime_type, data) = download_and_encode_s3_image(
|
||||
s3_object,
|
||||
client,
|
||||
workspace_id,
|
||||
)
|
||||
.await?;
|
||||
parts.push(GeminiPart::InlineData {
|
||||
inline_data: GeminiInlineData { mime_type, data },
|
||||
});
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
Ok(parts)
|
||||
}
|
||||
|
||||
/// Find function name by tool call ID from previous messages
|
||||
fn find_function_name_by_id(&self, messages: &[OpenAIMessage], tool_call_id: &str) -> String {
|
||||
for msg in messages {
|
||||
if let Some(tool_calls) = &msg.tool_calls {
|
||||
for tc in tool_calls {
|
||||
if tc.id == tool_call_id {
|
||||
return tc.function.name.clone();
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
"unknown_function".to_string()
|
||||
}
|
||||
|
||||
/// Convert OpenAI tools to Gemini format
|
||||
/// Convert OpenAI tool definitions to Gemini format.
|
||||
///
|
||||
/// Sanitizes each tool's JSON schema for Google compatibility before delegating
|
||||
/// to the shared [`openai_tools_to_gemini`] function.
|
||||
fn convert_tools_to_gemini(
|
||||
&self,
|
||||
tools: Option<&[ToolDef]>,
|
||||
has_websearch: bool,
|
||||
) -> Option<Vec<GeminiTool>> {
|
||||
let mut gemini_tools = Vec::new();
|
||||
|
||||
// Add function declarations
|
||||
if let Some(tool_defs) = tools {
|
||||
let declarations: Vec<GeminiFunctionDeclaration> = tool_defs
|
||||
.iter()
|
||||
.filter_map(|t| {
|
||||
// Deserialize RawValue into OpenAPISchema, sanitize, then use
|
||||
let mut schema: OpenAPISchema =
|
||||
serde_json::from_str(t.function.parameters.get()).ok()?;
|
||||
schema.sanitize_for_google();
|
||||
|
||||
Some(GeminiFunctionDeclaration {
|
||||
name: t.function.name.clone(),
|
||||
description: t.function.description.clone(),
|
||||
parameters: schema,
|
||||
})
|
||||
})
|
||||
.collect();
|
||||
|
||||
if !declarations.is_empty() {
|
||||
gemini_tools.push(GeminiTool {
|
||||
function_declarations: Some(declarations),
|
||||
google_search: None,
|
||||
});
|
||||
let Some(tool_defs) = tools else {
|
||||
if has_websearch {
|
||||
return Some(vec![GeminiTool {
|
||||
function_declarations: None,
|
||||
google_search: Some(serde_json::json!({})),
|
||||
}]);
|
||||
}
|
||||
}
|
||||
return None;
|
||||
};
|
||||
|
||||
// Add Google Search tool if enabled
|
||||
if has_websearch {
|
||||
gemini_tools.push(GeminiTool {
|
||||
function_declarations: None,
|
||||
google_search: Some(serde_json::json!({})),
|
||||
});
|
||||
}
|
||||
let tool_params: Vec<serde_json::Value> = tool_defs
|
||||
.iter()
|
||||
.map(|t| {
|
||||
let mut schema: OpenAPISchema =
|
||||
serde_json::from_str(t.function.parameters.get()).unwrap_or_default();
|
||||
schema.sanitize_for_google();
|
||||
serde_json::to_value(&schema).unwrap_or_default()
|
||||
})
|
||||
.collect();
|
||||
|
||||
if gemini_tools.is_empty() {
|
||||
None
|
||||
} else {
|
||||
Some(gemini_tools)
|
||||
}
|
||||
openai_tools_to_gemini(tool_defs, &tool_params, has_websearch)
|
||||
}
|
||||
|
||||
/// Build generation config for structured output and other settings
|
||||
fn build_generation_config(
|
||||
&self,
|
||||
args: &BuildRequestArgs<'_>,
|
||||
) -> Option<GeminiGenerationConfig> {
|
||||
fn build_generation_config(&self, args: &BuildRequestArgs<'_>) -> Option<GeminiGenerationConfig> {
|
||||
let has_output_schema = args
|
||||
.output_schema
|
||||
.and_then(|s| s.properties.as_ref())
|
||||
@@ -529,15 +142,11 @@ impl GoogleAIQueryBuilder {
|
||||
let (response_mime_type, response_schema) = if has_output_schema {
|
||||
let mut schema = args.output_schema.unwrap().clone();
|
||||
schema.sanitize_for_google();
|
||||
(
|
||||
Some("application/json".to_string()),
|
||||
serde_json::to_value(&schema).ok(),
|
||||
)
|
||||
(Some("application/json".to_string()), serde_json::to_value(&schema).ok())
|
||||
} else {
|
||||
(None, None)
|
||||
};
|
||||
|
||||
// Only create config if there's something to configure
|
||||
if args.temperature.is_some() || args.max_tokens.is_some() || response_mime_type.is_some() {
|
||||
Some(GeminiGenerationConfig {
|
||||
temperature: args.temperature,
|
||||
@@ -554,7 +163,6 @@ impl GoogleAIQueryBuilder {
|
||||
#[async_trait]
|
||||
impl QueryBuilder for GoogleAIQueryBuilder {
|
||||
fn supports_tools_with_output_type(&self, output_type: &OutputType) -> bool {
|
||||
// Google AI supports tools only for text output
|
||||
matches!(output_type, OutputType::Text)
|
||||
}
|
||||
|
||||
@@ -578,7 +186,6 @@ impl QueryBuilder for GoogleAIQueryBuilder {
|
||||
Error::internal_err(format!("Failed to parse Gemini image response: {}", e))
|
||||
})?;
|
||||
|
||||
// First, check Gemini models (candidates -> content -> parts -> inline_data)
|
||||
let image_data_from_gemini = gemini_response.candidates.as_ref().and_then(|candidates| {
|
||||
candidates.iter().find_map(|candidate| {
|
||||
candidate
|
||||
@@ -589,13 +196,11 @@ impl QueryBuilder for GoogleAIQueryBuilder {
|
||||
})
|
||||
});
|
||||
|
||||
// Then, check Imagen models (predictions -> bytes_base64_encoded)
|
||||
let image_data_from_imagen = gemini_response
|
||||
.predictions
|
||||
.as_ref()
|
||||
.and_then(|predictions| predictions.first().map(|p| &p.bytes_base64_encoded));
|
||||
|
||||
// Image data, preferring Gemini first then Imagen models
|
||||
let image_data = image_data_from_gemini.or(image_data_from_imagen);
|
||||
|
||||
match image_data {
|
||||
@@ -627,7 +232,6 @@ impl QueryBuilder for GoogleAIQueryBuilder {
|
||||
..
|
||||
} = gemini_sse_parser;
|
||||
|
||||
// Send tool call arguments events for accumulated tool calls
|
||||
for tool_call in accumulated_tool_calls.values() {
|
||||
let event = StreamingEvent::ToolCallArguments {
|
||||
call_id: tool_call.id.clone(),
|
||||
@@ -637,7 +241,6 @@ impl QueryBuilder for GoogleAIQueryBuilder {
|
||||
stream_event_processor.send(event, &mut events_str).await?;
|
||||
}
|
||||
|
||||
// Convert Gemini usage metadata to TokenUsage
|
||||
let usage = gemini_usage.map(|u| {
|
||||
TokenUsage::new(
|
||||
u.prompt_token_count,
|
||||
@@ -647,11 +250,7 @@ impl QueryBuilder for GoogleAIQueryBuilder {
|
||||
});
|
||||
|
||||
Ok(ParsedResponse::Text {
|
||||
content: if accumulated_content.is_empty() {
|
||||
None
|
||||
} else {
|
||||
Some(accumulated_content)
|
||||
},
|
||||
content: if accumulated_content.is_empty() { None } else { Some(accumulated_content) },
|
||||
tool_calls: accumulated_tool_calls.into_values().collect(),
|
||||
events_str: Some(events_str),
|
||||
annotations,
|
||||
@@ -663,17 +262,11 @@ impl QueryBuilder for GoogleAIQueryBuilder {
|
||||
fn get_endpoint(&self, base_url: &str, model: &str, output_type: &OutputType) -> String {
|
||||
match output_type {
|
||||
OutputType::Text => {
|
||||
format!(
|
||||
"{}/models/{}:streamGenerateContent?alt=sse",
|
||||
base_url, model
|
||||
)
|
||||
format!("{}/models/{}:streamGenerateContent?alt=sse", base_url, model)
|
||||
}
|
||||
OutputType::Image => {
|
||||
let url_suffix = if model.contains("imagen") {
|
||||
"predict"
|
||||
} else {
|
||||
"generateContent"
|
||||
};
|
||||
let url_suffix =
|
||||
if model.contains("imagen") { "predict" } else { "generateContent" };
|
||||
format!("{}/models/{}:{}", base_url, model, url_suffix)
|
||||
}
|
||||
}
|
||||
@@ -685,7 +278,6 @@ impl QueryBuilder for GoogleAIQueryBuilder {
|
||||
_base_url: &str,
|
||||
_output_type: &OutputType,
|
||||
) -> Vec<(&'static str, String)> {
|
||||
// Native Gemini API always uses x-goog-api-key
|
||||
vec![("x-goog-api-key", api_key.to_string())]
|
||||
}
|
||||
}
|
||||
|
||||
@@ -5,15 +5,18 @@ use reqwest::Response;
|
||||
use serde::Deserialize;
|
||||
use serde_json;
|
||||
use tokio_stream::StreamExt;
|
||||
use windmill_common::{error::Error, utils::rd_string};
|
||||
use windmill_common::{
|
||||
ai_google::{parse_gemini_sse_event, GeminiUsageMetadata},
|
||||
ai_types::{ExtraContent, GoogleExtraContent, OpenAIFunction, OpenAIToolCall},
|
||||
error::Error,
|
||||
utils::rd_string,
|
||||
};
|
||||
|
||||
use crate::ai::{
|
||||
query_builder::StreamEventProcessor,
|
||||
types::{StreamingEvent, UrlCitation},
|
||||
};
|
||||
|
||||
use windmill_common::ai_types::{ExtraContent, GoogleExtraContent, OpenAIFunction, OpenAIToolCall};
|
||||
|
||||
#[derive(Deserialize)]
|
||||
pub struct OpenAIChoiceDeltaToolCallFunction {
|
||||
pub name: Option<String>,
|
||||
@@ -457,96 +460,19 @@ impl SSEParser for AnthropicSSEParser {
|
||||
// Gemini SSE Parser
|
||||
// ============================================================================
|
||||
|
||||
/// Gemini streaming response part - can be text or function call
|
||||
#[derive(Deserialize, Debug)]
|
||||
pub struct GeminiSSEPart {
|
||||
#[serde(default)]
|
||||
pub text: Option<String>,
|
||||
#[serde(rename = "functionCall")]
|
||||
pub function_call: Option<GeminiSSEFunctionCall>,
|
||||
/// Thought signature for Gemini 3+ models - required for function calling
|
||||
#[serde(rename = "thoughtSignature")]
|
||||
pub thought_signature: Option<String>,
|
||||
}
|
||||
|
||||
/// Function call in Gemini streaming response
|
||||
#[derive(Deserialize, Debug)]
|
||||
pub struct GeminiSSEFunctionCall {
|
||||
pub name: String,
|
||||
pub args: serde_json::Value,
|
||||
}
|
||||
|
||||
/// Content in Gemini streaming candidate
|
||||
#[derive(Deserialize, Debug)]
|
||||
pub struct GeminiSSEContent {
|
||||
pub parts: Option<Vec<GeminiSSEPart>>,
|
||||
}
|
||||
|
||||
/// Web reference in Gemini grounding chunk
|
||||
#[derive(Deserialize, Debug)]
|
||||
pub struct GeminiGroundingChunkWeb {
|
||||
pub uri: String,
|
||||
#[serde(default)]
|
||||
pub title: Option<String>,
|
||||
}
|
||||
|
||||
/// Grounding chunk from Gemini web search
|
||||
#[derive(Deserialize, Debug)]
|
||||
pub struct GeminiGroundingChunk {
|
||||
pub web: Option<GeminiGroundingChunkWeb>,
|
||||
}
|
||||
|
||||
/// Grounding metadata from Gemini web search
|
||||
#[derive(Deserialize, Debug)]
|
||||
pub struct GeminiGroundingMetadata {
|
||||
#[serde(rename = "groundingChunks", default)]
|
||||
pub grounding_chunks: Vec<GeminiGroundingChunk>,
|
||||
#[serde(rename = "webSearchQueries", default)]
|
||||
pub web_search_queries: Vec<String>,
|
||||
}
|
||||
|
||||
/// Candidate in Gemini streaming response
|
||||
#[derive(Deserialize, Debug)]
|
||||
pub struct GeminiSSECandidate {
|
||||
pub content: Option<GeminiSSEContent>,
|
||||
#[serde(rename = "finishReason")]
|
||||
#[allow(dead_code)]
|
||||
pub finish_reason: Option<String>,
|
||||
#[serde(rename = "groundingMetadata")]
|
||||
pub grounding_metadata: Option<GeminiGroundingMetadata>,
|
||||
}
|
||||
|
||||
/// Gemini usage metadata from SSE response
|
||||
#[derive(Deserialize, Debug, Clone)]
|
||||
pub struct GeminiUsageMetadata {
|
||||
#[serde(rename = "promptTokenCount", default)]
|
||||
pub prompt_token_count: Option<i32>,
|
||||
#[serde(rename = "candidatesTokenCount", default)]
|
||||
pub candidates_token_count: Option<i32>,
|
||||
#[serde(rename = "totalTokenCount", default)]
|
||||
pub total_token_count: Option<i32>,
|
||||
}
|
||||
|
||||
/// Gemini SSE event structure
|
||||
#[derive(Deserialize, Debug)]
|
||||
pub struct GeminiSSEEvent {
|
||||
pub candidates: Option<Vec<GeminiSSECandidate>>,
|
||||
#[serde(rename = "usageMetadata")]
|
||||
pub usage_metadata: Option<GeminiUsageMetadata>,
|
||||
}
|
||||
|
||||
/// Gemini SSE Parser for streaming responses
|
||||
/// Accumulates Gemini streaming events and converts them into the worker's
|
||||
/// internal [`OpenAIToolCall`] / [`StreamingEvent`] representation.
|
||||
///
|
||||
/// The actual SSE parsing is delegated to [`parse_gemini_sse_event`] from
|
||||
/// `windmill_common::ai_google` so the logic can be shared with the API proxy.
|
||||
pub struct GeminiSSEParser {
|
||||
pub accumulated_content: String,
|
||||
pub accumulated_tool_calls: HashMap<i64, OpenAIToolCall>,
|
||||
pub events_str: String,
|
||||
pub stream_event_processor: StreamEventProcessor,
|
||||
tool_call_index: i64,
|
||||
/// Collected URL citation annotations from web search
|
||||
pub annotations: Vec<UrlCitation>,
|
||||
/// Whether web search was used in this response
|
||||
pub used_websearch: bool,
|
||||
/// Token usage from usageMetadata
|
||||
pub usage: Option<GeminiUsageMetadata>,
|
||||
}
|
||||
|
||||
@@ -567,101 +493,57 @@ impl GeminiSSEParser {
|
||||
|
||||
impl SSEParser for GeminiSSEParser {
|
||||
async fn parse_event_data(&mut self, data: &str) -> Result<(), Error> {
|
||||
let event: Option<GeminiSSEEvent> = serde_json::from_str(data)
|
||||
.inspect_err(|e| {
|
||||
tracing::error!("Failed to parse SSE as a Gemini event {}: {}", data, e);
|
||||
})
|
||||
.ok();
|
||||
let Some(parsed) = parse_gemini_sse_event(data)? else {
|
||||
return Ok(());
|
||||
};
|
||||
|
||||
if let Some(event) = event {
|
||||
if let Some(candidates) = event.candidates {
|
||||
for candidate in candidates {
|
||||
if let Some(content) = candidate.content {
|
||||
if let Some(parts) = content.parts {
|
||||
for part in parts {
|
||||
// Handle text content
|
||||
if let Some(text) = part.text {
|
||||
if !text.is_empty() {
|
||||
self.accumulated_content.push_str(&text);
|
||||
let event = StreamingEvent::TokenDelta { content: text };
|
||||
self.stream_event_processor
|
||||
.send(event, &mut self.events_str)
|
||||
.await?;
|
||||
}
|
||||
}
|
||||
if let Some(text) = parsed.text {
|
||||
self.accumulated_content.push_str(&text);
|
||||
self.stream_event_processor
|
||||
.send(StreamingEvent::TokenDelta { content: text }, &mut self.events_str)
|
||||
.await?;
|
||||
}
|
||||
|
||||
// Handle function calls
|
||||
if let Some(function_call) = part.function_call {
|
||||
let call_id = format!("call_{}", rd_string(24));
|
||||
let idx = self.tool_call_index;
|
||||
self.tool_call_index += 1;
|
||||
for tool_call in parsed.tool_calls {
|
||||
let call_id = format!("call_{}", rd_string(24));
|
||||
let idx = self.tool_call_index;
|
||||
self.tool_call_index += 1;
|
||||
|
||||
// Send tool call start event
|
||||
let event = StreamingEvent::ToolCall {
|
||||
call_id: call_id.clone(),
|
||||
function_name: function_call.name.clone(),
|
||||
};
|
||||
self.stream_event_processor
|
||||
.send(event, &mut self.events_str)
|
||||
.await?;
|
||||
self.stream_event_processor
|
||||
.send(
|
||||
StreamingEvent::ToolCall {
|
||||
call_id: call_id.clone(),
|
||||
function_name: tool_call.name.clone(),
|
||||
},
|
||||
&mut self.events_str,
|
||||
)
|
||||
.await?;
|
||||
|
||||
// Build extra_content with thought_signature if present
|
||||
let extra_content =
|
||||
part.thought_signature.map(|sig| ExtraContent {
|
||||
google: Some(GoogleExtraContent {
|
||||
thought_signature: Some(sig),
|
||||
}),
|
||||
});
|
||||
let extra_content = tool_call.thought_signature.map(|sig| ExtraContent {
|
||||
google: Some(GoogleExtraContent { thought_signature: Some(sig) }),
|
||||
});
|
||||
|
||||
// Store accumulated tool call
|
||||
self.accumulated_tool_calls.insert(
|
||||
idx,
|
||||
OpenAIToolCall {
|
||||
id: call_id,
|
||||
function: OpenAIFunction {
|
||||
name: function_call.name,
|
||||
arguments: serde_json::to_string(
|
||||
&function_call.args,
|
||||
)
|
||||
.unwrap_or_else(|_| "{}".to_string()),
|
||||
},
|
||||
r#type: "function".to_string(),
|
||||
extra_content,
|
||||
},
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
self.accumulated_tool_calls.insert(
|
||||
idx,
|
||||
OpenAIToolCall {
|
||||
id: call_id,
|
||||
function: OpenAIFunction {
|
||||
name: tool_call.name,
|
||||
arguments: serde_json::to_string(&tool_call.args)
|
||||
.unwrap_or_else(|_| "{}".to_string()),
|
||||
},
|
||||
r#type: "function".to_string(),
|
||||
extra_content,
|
||||
},
|
||||
);
|
||||
}
|
||||
|
||||
// Handle grounding metadata (web search results)
|
||||
if let Some(ref grounding_metadata) = candidate.grounding_metadata {
|
||||
// Set used_websearch if there are search queries or grounding chunks
|
||||
if !grounding_metadata.web_search_queries.is_empty()
|
||||
|| !grounding_metadata.grounding_chunks.is_empty()
|
||||
{
|
||||
self.used_websearch = true;
|
||||
}
|
||||
|
||||
// Extract citations from grounding chunks
|
||||
for chunk in &grounding_metadata.grounding_chunks {
|
||||
if let Some(ref web) = chunk.web {
|
||||
self.annotations.push(UrlCitation {
|
||||
start_index: 0, // Gemini doesn't provide character indices
|
||||
end_index: 0,
|
||||
url: web.uri.clone(),
|
||||
title: web.title.clone(),
|
||||
});
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Extract usage metadata
|
||||
if let Some(usage_metadata) = event.usage_metadata {
|
||||
self.usage = Some(usage_metadata);
|
||||
}
|
||||
self.annotations.extend(parsed.annotations);
|
||||
if parsed.used_websearch {
|
||||
self.used_websearch = true;
|
||||
}
|
||||
if let Some(usage) = parsed.usage {
|
||||
self.usage = Some(usage);
|
||||
}
|
||||
|
||||
Ok(())
|
||||
|
||||
@@ -731,16 +731,3 @@ pub fn extract_text_content(content: &OpenAIContent) -> String {
|
||||
.join(""),
|
||||
}
|
||||
}
|
||||
|
||||
/// Parse a data URL to extract media type and base64 data
|
||||
/// Format: data:mime_type;base64,data
|
||||
/// Returns (media_type, data) tuple if successful
|
||||
pub fn parse_data_url(url: &str) -> Option<(String, String)> {
|
||||
if !url.starts_with("data:") {
|
||||
return None;
|
||||
}
|
||||
let rest = url.strip_prefix("data:")?;
|
||||
let (header, data) = rest.split_once(",")?;
|
||||
let media_type = header.strip_suffix(";base64")?;
|
||||
Some((media_type.to_string(), data.to_string()))
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user