use serde::{Deserialize, Serialize};
use serde_json::Value;
use std::collections::HashMap;
use crate::enums::{
ApiSpec, AuthType, Behavior, DynamicRetrievalConfigMode, Environment, FunctionCallingMode,
HttpElementLocation, PhishBlockThreshold, Type,
};
#[derive(Debug, Clone, Serialize, Deserialize, Default)]
#[serde(rename_all = "camelCase")]
pub struct Tool {
#[serde(skip_serializing_if = "Option::is_none")]
pub retrieval: Option<Retrieval>,
#[serde(skip_serializing_if = "Option::is_none")]
pub computer_use: Option<ComputerUse>,
#[serde(skip_serializing_if = "Option::is_none")]
pub file_search: Option<FileSearch>,
#[serde(skip_serializing_if = "Option::is_none")]
pub code_execution: Option<CodeExecution>,
#[serde(skip_serializing_if = "Option::is_none")]
pub enterprise_web_search: Option<EnterpriseWebSearch>,
#[serde(skip_serializing_if = "Option::is_none")]
pub function_declarations: Option<Vec<FunctionDeclaration>>,
#[serde(skip_serializing_if = "Option::is_none")]
pub google_maps: Option<GoogleMaps>,
#[serde(skip_serializing_if = "Option::is_none")]
pub google_search: Option<GoogleSearch>,
#[serde(skip_serializing_if = "Option::is_none")]
pub google_search_retrieval: Option<GoogleSearchRetrieval>,
#[serde(skip_serializing_if = "Option::is_none")]
pub url_context: Option<UrlContext>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct FunctionDeclaration {
pub name: String,
#[serde(skip_serializing_if = "Option::is_none")]
pub description: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub parameters: Option<Schema>,
#[serde(skip_serializing_if = "Option::is_none")]
pub parameters_json_schema: Option<Value>,
#[serde(skip_serializing_if = "Option::is_none")]
pub response: Option<Schema>,
#[serde(skip_serializing_if = "Option::is_none")]
pub response_json_schema: Option<Value>,
#[serde(skip_serializing_if = "Option::is_none")]
pub behavior: Option<Behavior>,
}
#[derive(Debug, Clone, Serialize, Deserialize, Default)]
#[serde(rename_all = "camelCase")]
pub struct GoogleSearch {
#[serde(skip_serializing_if = "Option::is_none")]
pub exclude_domains: Option<Vec<String>>,
#[serde(skip_serializing_if = "Option::is_none")]
pub blocking_confidence: Option<PhishBlockThreshold>,
#[serde(skip_serializing_if = "Option::is_none")]
pub time_range_filter: Option<Interval>,
}
#[derive(Debug, Clone, Serialize, Deserialize, Default)]
#[serde(rename_all = "camelCase")]
pub struct EnterpriseWebSearch {
#[serde(skip_serializing_if = "Option::is_none")]
pub exclude_domains: Option<Vec<String>>,
#[serde(skip_serializing_if = "Option::is_none")]
pub blocking_confidence: Option<PhishBlockThreshold>,
}
#[derive(Debug, Clone, Serialize, Deserialize, Default)]
#[serde(rename_all = "camelCase")]
pub struct CodeExecution {}
#[derive(Debug, Clone, Serialize, Deserialize, Default)]
#[serde(rename_all = "camelCase")]
pub struct UrlContext {}
#[derive(Debug, Clone, Serialize, Deserialize, Default)]
#[serde(rename_all = "camelCase")]
pub struct ComputerUse {
#[serde(skip_serializing_if = "Option::is_none")]
pub environment: Option<Environment>,
#[serde(skip_serializing_if = "Option::is_none")]
pub excluded_predefined_functions: Option<Vec<String>>,
}
#[derive(Debug, Clone, Serialize, Deserialize, Default)]
#[serde(rename_all = "camelCase")]
pub struct GoogleMaps {
#[serde(skip_serializing_if = "Option::is_none")]
pub auth_config: Option<AuthConfig>,
#[serde(skip_serializing_if = "Option::is_none")]
pub enable_widget: Option<bool>,
}
#[derive(Debug, Clone, Serialize, Deserialize, Default)]
#[serde(rename_all = "camelCase")]
pub struct GoogleSearchRetrieval {
#[serde(skip_serializing_if = "Option::is_none")]
pub dynamic_retrieval_config: Option<DynamicRetrievalConfig>,
}
#[derive(Debug, Clone, Serialize, Deserialize, Default)]
#[serde(rename_all = "camelCase")]
pub struct FileSearch {
#[serde(skip_serializing_if = "Option::is_none")]
pub file_search_store_names: Option<Vec<String>>,
#[serde(skip_serializing_if = "Option::is_none")]
pub top_k: Option<i32>,
#[serde(skip_serializing_if = "Option::is_none")]
pub metadata_filter: Option<String>,
}
#[derive(Debug, Clone, Serialize, Deserialize, Default)]
#[serde(rename_all = "camelCase")]
pub struct Interval {
#[serde(skip_serializing_if = "Option::is_none")]
pub end_time: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub start_time: Option<String>,
}
#[derive(Debug, Clone, Serialize, Deserialize, Default)]
#[serde(rename_all = "camelCase")]
pub struct ApiKeyConfig {
#[serde(skip_serializing_if = "Option::is_none")]
pub api_key_secret: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub api_key_string: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub http_element_location: Option<HttpElementLocation>,
#[serde(skip_serializing_if = "Option::is_none")]
pub name: Option<String>,
}
#[derive(Debug, Clone, Serialize, Deserialize, Default)]
#[serde(rename_all = "camelCase")]
pub struct ApiAuthApiKeyConfig {
#[serde(skip_serializing_if = "Option::is_none")]
pub api_key_secret_version: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub api_key_string: Option<String>,
}
#[derive(Debug, Clone, Serialize, Deserialize, Default)]
#[serde(rename_all = "camelCase")]
pub struct ApiAuth {
#[serde(skip_serializing_if = "Option::is_none")]
pub api_key_config: Option<ApiAuthApiKeyConfig>,
}
#[derive(Debug, Clone, Serialize, Deserialize, Default)]
#[serde(rename_all = "camelCase")]
pub struct AuthConfigGoogleServiceAccountConfig {
#[serde(skip_serializing_if = "Option::is_none")]
pub service_account: Option<String>,
}
#[derive(Debug, Clone, Serialize, Deserialize, Default)]
#[serde(rename_all = "camelCase")]
pub struct AuthConfigHttpBasicAuthConfig {
#[serde(skip_serializing_if = "Option::is_none")]
pub credential_secret: Option<String>,
}
#[derive(Debug, Clone, Serialize, Deserialize, Default)]
#[serde(rename_all = "camelCase")]
pub struct AuthConfigOauthConfig {
#[serde(skip_serializing_if = "Option::is_none")]
pub access_token: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub service_account: Option<String>,
}
#[derive(Debug, Clone, Serialize, Deserialize, Default)]
#[serde(rename_all = "camelCase")]
pub struct AuthConfigOidcConfig {
#[serde(skip_serializing_if = "Option::is_none")]
pub id_token: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub service_account: Option<String>,
}
#[derive(Debug, Clone, Serialize, Deserialize, Default)]
#[serde(rename_all = "camelCase")]
pub struct AuthConfig {
#[serde(skip_serializing_if = "Option::is_none")]
pub api_key_config: Option<ApiKeyConfig>,
#[serde(skip_serializing_if = "Option::is_none")]
pub auth_type: Option<AuthType>,
#[serde(skip_serializing_if = "Option::is_none")]
pub google_service_account_config: Option<AuthConfigGoogleServiceAccountConfig>,
#[serde(skip_serializing_if = "Option::is_none")]
pub http_basic_auth_config: Option<AuthConfigHttpBasicAuthConfig>,
#[serde(skip_serializing_if = "Option::is_none")]
pub oauth_config: Option<AuthConfigOauthConfig>,
#[serde(skip_serializing_if = "Option::is_none")]
pub oidc_config: Option<AuthConfigOidcConfig>,
}
#[derive(Debug, Clone, Serialize, Deserialize, Default)]
#[serde(rename_all = "camelCase")]
pub struct ExternalApiElasticSearchParams {
#[serde(skip_serializing_if = "Option::is_none")]
pub index: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub num_hits: Option<i32>,
#[serde(skip_serializing_if = "Option::is_none")]
pub search_template: Option<String>,
}
#[derive(Debug, Clone, Serialize, Deserialize, Default)]
#[serde(rename_all = "camelCase")]
pub struct ExternalApiSimpleSearchParams {}
#[derive(Debug, Clone, Serialize, Deserialize, Default)]
#[serde(rename_all = "camelCase")]
pub struct ExternalApi {
#[serde(skip_serializing_if = "Option::is_none")]
pub api_auth: Option<ApiAuth>,
#[serde(skip_serializing_if = "Option::is_none")]
pub api_spec: Option<ApiSpec>,
#[serde(skip_serializing_if = "Option::is_none")]
pub auth_config: Option<AuthConfig>,
#[serde(skip_serializing_if = "Option::is_none")]
pub elastic_search_params: Option<ExternalApiElasticSearchParams>,
#[serde(skip_serializing_if = "Option::is_none")]
pub endpoint: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub simple_search_params: Option<ExternalApiSimpleSearchParams>,
}
#[derive(Debug, Clone, Serialize, Deserialize, Default)]
#[serde(rename_all = "camelCase")]
pub struct VertexAiSearchDataStoreSpec {
#[serde(skip_serializing_if = "Option::is_none")]
pub data_store: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub filter: Option<String>,
}
#[derive(Debug, Clone, Serialize, Deserialize, Default)]
#[serde(rename_all = "camelCase")]
pub struct VertexAiSearch {
#[serde(skip_serializing_if = "Option::is_none")]
pub data_store_specs: Option<Vec<VertexAiSearchDataStoreSpec>>,
#[serde(skip_serializing_if = "Option::is_none")]
pub datastore: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub engine: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub filter: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub max_results: Option<i32>,
}
#[derive(Debug, Clone, Serialize, Deserialize, Default)]
#[serde(rename_all = "camelCase")]
pub struct VertexRagStoreRagResource {
#[serde(skip_serializing_if = "Option::is_none")]
pub rag_corpus: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub rag_file_ids: Option<Vec<String>>,
}
#[derive(Debug, Clone, Serialize, Deserialize, Default)]
#[serde(rename_all = "camelCase")]
pub struct RagRetrievalConfigFilter {
#[serde(skip_serializing_if = "Option::is_none")]
pub metadata_filter: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub vector_distance_threshold: Option<f64>,
#[serde(skip_serializing_if = "Option::is_none")]
pub vector_similarity_threshold: Option<f64>,
}
#[derive(Debug, Clone, Serialize, Deserialize, Default)]
#[serde(rename_all = "camelCase")]
pub struct RagRetrievalConfigHybridSearch {
#[serde(skip_serializing_if = "Option::is_none")]
pub alpha: Option<f32>,
}
#[derive(Debug, Clone, Serialize, Deserialize, Default)]
#[serde(rename_all = "camelCase")]
pub struct RagRetrievalConfigRankingLlmRanker {
#[serde(skip_serializing_if = "Option::is_none")]
pub model_name: Option<String>,
}
#[derive(Debug, Clone, Serialize, Deserialize, Default)]
#[serde(rename_all = "camelCase")]
pub struct RagRetrievalConfigRankingRankService {
#[serde(skip_serializing_if = "Option::is_none")]
pub model_name: Option<String>,
}
#[derive(Debug, Clone, Serialize, Deserialize, Default)]
#[serde(rename_all = "camelCase")]
pub struct RagRetrievalConfigRanking {
#[serde(skip_serializing_if = "Option::is_none")]
pub llm_ranker: Option<RagRetrievalConfigRankingLlmRanker>,
#[serde(skip_serializing_if = "Option::is_none")]
pub rank_service: Option<RagRetrievalConfigRankingRankService>,
}
#[derive(Debug, Clone, Serialize, Deserialize, Default)]
#[serde(rename_all = "camelCase")]
pub struct RagRetrievalConfig {
#[serde(skip_serializing_if = "Option::is_none")]
pub filter: Option<RagRetrievalConfigFilter>,
#[serde(skip_serializing_if = "Option::is_none")]
pub hybrid_search: Option<RagRetrievalConfigHybridSearch>,
#[serde(skip_serializing_if = "Option::is_none")]
pub ranking: Option<RagRetrievalConfigRanking>,
#[serde(skip_serializing_if = "Option::is_none")]
pub top_k: Option<i32>,
}
#[derive(Debug, Clone, Serialize, Deserialize, Default)]
#[serde(rename_all = "camelCase")]
pub struct VertexRagStore {
#[serde(skip_serializing_if = "Option::is_none")]
pub rag_corpora: Option<Vec<String>>,
#[serde(skip_serializing_if = "Option::is_none")]
pub rag_resources: Option<Vec<VertexRagStoreRagResource>>,
#[serde(skip_serializing_if = "Option::is_none")]
pub rag_retrieval_config: Option<RagRetrievalConfig>,
#[serde(skip_serializing_if = "Option::is_none")]
pub similarity_top_k: Option<i32>,
#[serde(skip_serializing_if = "Option::is_none")]
pub store_context: Option<bool>,
#[serde(skip_serializing_if = "Option::is_none")]
pub vector_distance_threshold: Option<f64>,
}
#[derive(Debug, Clone, Serialize, Deserialize, Default)]
#[serde(rename_all = "camelCase")]
pub struct Retrieval {
#[serde(skip_serializing_if = "Option::is_none")]
pub disable_attribution: Option<bool>,
#[serde(skip_serializing_if = "Option::is_none")]
pub external_api: Option<ExternalApi>,
#[serde(skip_serializing_if = "Option::is_none")]
pub vertex_ai_search: Option<VertexAiSearch>,
#[serde(skip_serializing_if = "Option::is_none")]
pub vertex_rag_store: Option<VertexRagStore>,
}
#[derive(Debug, Clone, Serialize, Deserialize, Default)]
#[serde(rename_all = "camelCase")]
pub struct DynamicRetrievalConfig {
#[serde(skip_serializing_if = "Option::is_none")]
pub dynamic_threshold: Option<f32>,
#[serde(skip_serializing_if = "Option::is_none")]
pub mode: Option<DynamicRetrievalConfigMode>,
}
#[derive(Debug, Clone, Serialize, Deserialize, Default)]
#[serde(rename_all = "camelCase")]
pub struct ToolConfig {
#[serde(skip_serializing_if = "Option::is_none")]
pub function_calling_config: Option<FunctionCallingConfig>,
#[serde(skip_serializing_if = "Option::is_none")]
pub retrieval_config: Option<RetrievalConfig>,
}
#[derive(Debug, Clone, Serialize, Deserialize, Default)]
#[serde(rename_all = "camelCase")]
pub struct FunctionCallingConfig {
#[serde(skip_serializing_if = "Option::is_none")]
pub allowed_function_names: Option<Vec<String>>,
#[serde(skip_serializing_if = "Option::is_none")]
pub mode: Option<FunctionCallingMode>,
#[serde(skip_serializing_if = "Option::is_none")]
pub stream_function_call_arguments: Option<bool>,
}
#[derive(Debug, Clone, Serialize, Deserialize, Default)]
#[serde(rename_all = "camelCase")]
pub struct LatLng {
#[serde(skip_serializing_if = "Option::is_none")]
pub latitude: Option<f64>,
#[serde(skip_serializing_if = "Option::is_none")]
pub longitude: Option<f64>,
}
#[derive(Debug, Clone, Serialize, Deserialize, Default)]
#[serde(rename_all = "camelCase")]
pub struct RetrievalConfig {
#[serde(skip_serializing_if = "Option::is_none")]
pub lat_lng: Option<LatLng>,
#[serde(skip_serializing_if = "Option::is_none")]
pub language_code: Option<String>,
}
#[derive(Debug, Clone, Serialize, Deserialize, Default)]
#[serde(rename_all = "camelCase")]
pub struct Schema {
#[serde(skip_serializing_if = "Option::is_none")]
pub any_of: Option<Vec<Schema>>,
#[serde(skip_serializing_if = "Option::is_none")]
pub default: Option<Value>,
#[serde(skip_serializing_if = "Option::is_none")]
pub description: Option<String>,
#[serde(skip_serializing_if = "Option::is_none", rename = "enum")]
pub enum_values: Option<Vec<String>>,
#[serde(skip_serializing_if = "Option::is_none")]
pub example: Option<Value>,
#[serde(skip_serializing_if = "Option::is_none")]
pub format: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub items: Option<Box<Schema>>,
#[serde(skip_serializing_if = "Option::is_none")]
pub max_items: Option<i64>,
#[serde(skip_serializing_if = "Option::is_none")]
pub max_length: Option<i64>,
#[serde(skip_serializing_if = "Option::is_none")]
pub max_properties: Option<i64>,
#[serde(skip_serializing_if = "Option::is_none")]
pub maximum: Option<f64>,
#[serde(skip_serializing_if = "Option::is_none")]
pub min_items: Option<i64>,
#[serde(skip_serializing_if = "Option::is_none")]
pub min_length: Option<i64>,
#[serde(skip_serializing_if = "Option::is_none")]
pub min_properties: Option<i64>,
#[serde(skip_serializing_if = "Option::is_none")]
pub minimum: Option<f64>,
#[serde(skip_serializing_if = "Option::is_none")]
pub nullable: Option<bool>,
#[serde(skip_serializing_if = "Option::is_none")]
pub pattern: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub properties: Option<HashMap<String, Box<Schema>>>,
#[serde(skip_serializing_if = "Option::is_none")]
pub property_ordering: Option<Vec<String>>,
#[serde(skip_serializing_if = "Option::is_none")]
pub required: Option<Vec<String>>,
#[serde(skip_serializing_if = "Option::is_none")]
pub title: Option<String>,
#[serde(rename = "type", skip_serializing_if = "Option::is_none")]
pub ty: Option<Type>,
}
#[cfg(test)]
mod schema_builder_tests {
use super::*;
#[test]
fn test_tool_serialization() {
let tool = Tool {
google_maps: Some(GoogleMaps {
enable_widget: Some(true),
..GoogleMaps::default()
}),
..Tool::default()
};
let json = serde_json::to_value(&tool).unwrap();
assert_eq!(json["googleMaps"]["enableWidget"].as_bool(), Some(true));
}
}
impl Schema {
#[must_use]
pub fn object() -> SchemaBuilder {
SchemaBuilder::new(Type::Object)
}
#[must_use]
pub fn array() -> SchemaBuilder {
SchemaBuilder::new(Type::Array)
}
#[must_use]
pub fn string() -> Self {
Self {
ty: Some(Type::String),
..Default::default()
}
}
#[must_use]
pub fn integer() -> Self {
Self {
ty: Some(Type::Integer),
..Default::default()
}
}
#[must_use]
pub fn number() -> Self {
Self {
ty: Some(Type::Number),
..Default::default()
}
}
#[must_use]
pub fn boolean() -> Self {
Self {
ty: Some(Type::Boolean),
..Default::default()
}
}
}
pub struct SchemaBuilder {
schema: Schema,
}
impl SchemaBuilder {
#[must_use]
pub fn new(ty: Type) -> Self {
Self {
schema: Schema {
ty: Some(ty),
..Default::default()
},
}
}
#[must_use]
pub fn description(mut self, description: impl Into<String>) -> Self {
self.schema.description = Some(description.into());
self
}
#[must_use]
pub fn property(mut self, name: impl Into<String>, schema: Schema) -> Self {
let properties = self.schema.properties.get_or_insert_with(HashMap::new);
properties.insert(name.into(), Box::new(schema));
self
}
#[must_use]
pub fn required(mut self, name: impl Into<String>) -> Self {
let required = self.schema.required.get_or_insert_with(Vec::new);
required.push(name.into());
self
}
#[must_use]
pub fn items(mut self, schema: Schema) -> Self {
self.schema.items = Some(Box::new(schema));
self
}
#[must_use]
pub fn enum_values(mut self, values: Vec<String>) -> Self {
self.schema.enum_values = Some(values);
self
}
#[must_use]
pub fn build(self) -> Schema {
self.schema
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn schema_builder_object() {
let schema = Schema::object()
.property("name", Schema::string())
.required("name")
.build();
assert_eq!(schema.ty, Some(Type::Object));
assert!(schema.properties.unwrap().contains_key("name"));
}
#[test]
fn schema_builder_array_and_enum() {
let schema = Schema::array()
.items(Schema::string())
.enum_values(vec!["a".into(), "b".into()])
.build();
assert_eq!(schema.ty, Some(Type::Array));
assert_eq!(schema.items.unwrap().ty, Some(Type::String));
assert_eq!(
schema.enum_values.unwrap(),
vec!["a".to_string(), "b".to_string()]
);
}
#[test]
fn tool_function_declaration_serialization() {
let declaration = FunctionDeclaration {
name: "lookup".to_string(),
description: Some("search".to_string()),
parameters: Some(Schema::object().property("q", Schema::string()).build()),
parameters_json_schema: None,
response: Some(Schema::string()),
response_json_schema: None,
behavior: None,
};
let tool = Tool {
function_declarations: Some(vec![declaration]),
..Default::default()
};
let json = serde_json::to_value(&tool).unwrap();
assert!(json["functionDeclarations"].is_array());
assert_eq!(
json["functionDeclarations"][0]["name"].as_str(),
Some("lookup")
);
}
}