use serde::{Deserialize, Serialize};
use std::collections::HashMap;
#[derive(Clone, Debug)]
#[non_exhaustive]
pub enum Tool {
Function {
name: String,
description: String,
parameters: FunctionParameters,
},
GoogleSearch {
search_types: Option<Vec<SearchType>>,
},
GoogleMaps {
enable_widget: Option<bool>,
latitude: Option<f64>,
longitude: Option<f64>,
},
CodeExecution,
UrlContext,
ComputerUse {
environment: String,
excluded_predefined_functions: Vec<String>,
enable_prompt_injection_detection: Option<bool>,
disabled_safety_policies: Vec<String>,
},
McpServer {
name: String,
url: String,
allowed_tools: Option<Vec<AllowedTools>>,
headers: Option<HashMap<String, String>>,
},
FileSearch {
store_names: Vec<String>,
top_k: Option<i32>,
metadata_filter: Option<String>,
},
Retrieval {
retrieval_types: Option<Vec<RetrievalType>>,
vertex_ai_search_config: Option<VertexAiSearchConfig>,
exa_ai_search_config: Option<ExaAiSearchConfig>,
parallel_ai_search_config: Option<ParallelAiSearchConfig>,
rag_store_config: Option<Box<RagStoreConfig>>,
},
Unknown {
tool_type: String,
data: serde_json::Value,
},
}
impl Serialize for Tool {
fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
where
S: serde::Serializer,
{
use serde::ser::SerializeMap;
match self {
Self::Function {
name,
description,
parameters,
} => {
let mut map = serializer.serialize_map(None)?;
map.serialize_entry("type", "function")?;
map.serialize_entry("name", name)?;
map.serialize_entry("description", description)?;
map.serialize_entry("parameters", parameters)?;
map.end()
}
Self::GoogleSearch { search_types } => {
let mut map = serializer.serialize_map(None)?;
map.serialize_entry("type", "google_search")?;
if let Some(types) = search_types
&& !types.is_empty()
{
map.serialize_entry("search_types", types)?;
}
map.end()
}
Self::GoogleMaps {
enable_widget,
latitude,
longitude,
} => {
let mut map = serializer.serialize_map(None)?;
map.serialize_entry("type", "google_maps")?;
if let Some(ew) = enable_widget {
map.serialize_entry("enable_widget", ew)?;
}
if let Some(lat) = latitude {
map.serialize_entry("latitude", lat)?;
}
if let Some(lng) = longitude {
map.serialize_entry("longitude", lng)?;
}
map.end()
}
Self::CodeExecution => {
let mut map = serializer.serialize_map(None)?;
map.serialize_entry("type", "code_execution")?;
map.end()
}
Self::UrlContext => {
let mut map = serializer.serialize_map(None)?;
map.serialize_entry("type", "url_context")?;
map.end()
}
Self::ComputerUse {
environment,
excluded_predefined_functions,
enable_prompt_injection_detection,
disabled_safety_policies,
} => {
let mut map = serializer.serialize_map(None)?;
map.serialize_entry("type", "computer_use")?;
map.serialize_entry("environment", environment)?;
if !excluded_predefined_functions.is_empty() {
map.serialize_entry(
"excluded_predefined_functions",
excluded_predefined_functions,
)?;
}
if let Some(detect) = enable_prompt_injection_detection {
map.serialize_entry("enable_prompt_injection_detection", detect)?;
}
if !disabled_safety_policies.is_empty() {
map.serialize_entry("disabled_safety_policies", disabled_safety_policies)?;
}
map.end()
}
Self::McpServer {
name,
url,
allowed_tools,
headers,
} => {
let mut map = serializer.serialize_map(None)?;
map.serialize_entry("type", "mcp_server")?;
map.serialize_entry("name", name)?;
map.serialize_entry("url", url)?;
if let Some(tools) = allowed_tools
&& !tools.is_empty()
{
map.serialize_entry("allowed_tools", tools)?;
}
if let Some(hdrs) = headers
&& !hdrs.is_empty()
{
map.serialize_entry("headers", hdrs)?;
}
map.end()
}
Self::FileSearch {
store_names,
top_k,
metadata_filter,
} => {
let mut map = serializer.serialize_map(None)?;
map.serialize_entry("type", "file_search")?;
map.serialize_entry("file_search_store_names", store_names)?;
if let Some(k) = top_k {
map.serialize_entry("top_k", k)?;
}
if let Some(filter) = metadata_filter {
map.serialize_entry("metadata_filter", filter)?;
}
map.end()
}
Self::Retrieval {
retrieval_types,
vertex_ai_search_config,
exa_ai_search_config,
parallel_ai_search_config,
rag_store_config,
} => {
let mut map = serializer.serialize_map(None)?;
map.serialize_entry("type", "retrieval")?;
if let Some(types) = retrieval_types
&& !types.is_empty()
{
map.serialize_entry("retrieval_types", types)?;
}
if let Some(config) = vertex_ai_search_config {
map.serialize_entry("vertex_ai_search_config", config)?;
}
if let Some(config) = exa_ai_search_config {
map.serialize_entry("exa_ai_search_config", config)?;
}
if let Some(config) = parallel_ai_search_config {
map.serialize_entry("parallel_ai_search_config", config)?;
}
if let Some(config) = rag_store_config {
map.serialize_entry("rag_store_config", config)?;
}
map.end()
}
Self::Unknown { tool_type, data } => {
let mut map = serializer.serialize_map(None)?;
map.serialize_entry("type", tool_type)?;
if let serde_json::Value::Object(obj) = data {
for (key, value) in obj {
if key != "type" {
map.serialize_entry(key, value)?;
}
}
} else if !data.is_null() {
map.serialize_entry("data", data)?;
}
map.end()
}
}
}
}
impl<'de> Deserialize<'de> for Tool {
fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
where
D: serde::Deserializer<'de>,
{
let value = serde_json::Value::deserialize(deserializer)?;
#[derive(Deserialize)]
#[serde(tag = "type")]
enum KnownTool {
#[serde(rename = "function")]
Function {
name: String,
description: String,
parameters: FunctionParameters,
},
#[serde(rename = "google_search")]
GoogleSearch {
#[serde(default)]
search_types: Option<Vec<SearchType>>,
},
#[serde(rename = "google_maps")]
GoogleMaps {
#[serde(default)]
enable_widget: Option<bool>,
#[serde(default)]
latitude: Option<f64>,
#[serde(default)]
longitude: Option<f64>,
},
#[serde(rename = "code_execution")]
CodeExecution,
#[serde(rename = "url_context")]
UrlContext,
#[serde(rename = "computer_use")]
ComputerUse {
environment: String,
#[serde(default, alias = "excludedPredefinedFunctions")]
excluded_predefined_functions: Vec<String>,
#[serde(default)]
enable_prompt_injection_detection: Option<bool>,
#[serde(default)]
disabled_safety_policies: Vec<String>,
},
#[serde(rename = "mcp_server")]
McpServer {
name: String,
url: String,
#[serde(default)]
allowed_tools: Option<Vec<AllowedTools>>,
#[serde(default)]
headers: Option<HashMap<String, String>>,
},
#[serde(rename = "file_search")]
FileSearch {
#[serde(rename = "file_search_store_names")]
store_names: Vec<String>,
#[serde(default)]
top_k: Option<i32>,
#[serde(default)]
metadata_filter: Option<String>,
},
#[serde(rename = "retrieval")]
Retrieval {
#[serde(default)]
retrieval_types: Option<Vec<RetrievalType>>,
#[serde(default)]
vertex_ai_search_config: Option<VertexAiSearchConfig>,
#[serde(default)]
exa_ai_search_config: Option<ExaAiSearchConfig>,
#[serde(default)]
parallel_ai_search_config: Option<ParallelAiSearchConfig>,
#[serde(default)]
rag_store_config: Option<Box<RagStoreConfig>>,
},
}
match serde_json::from_value::<KnownTool>(value.clone()) {
Ok(known) => Ok(match known {
KnownTool::Function {
name,
description,
parameters,
} => Tool::Function {
name,
description,
parameters,
},
KnownTool::GoogleSearch { search_types } => Tool::GoogleSearch { search_types },
KnownTool::GoogleMaps {
enable_widget,
latitude,
longitude,
} => Tool::GoogleMaps {
enable_widget,
latitude,
longitude,
},
KnownTool::CodeExecution => Tool::CodeExecution,
KnownTool::UrlContext => Tool::UrlContext,
KnownTool::ComputerUse {
environment,
excluded_predefined_functions,
enable_prompt_injection_detection,
disabled_safety_policies,
} => Tool::ComputerUse {
environment,
excluded_predefined_functions,
enable_prompt_injection_detection,
disabled_safety_policies,
},
KnownTool::McpServer {
name,
url,
allowed_tools,
headers,
} => Tool::McpServer {
name,
url,
allowed_tools,
headers,
},
KnownTool::FileSearch {
store_names,
top_k,
metadata_filter,
} => Tool::FileSearch {
store_names,
top_k,
metadata_filter,
},
KnownTool::Retrieval {
retrieval_types,
vertex_ai_search_config,
exa_ai_search_config,
parallel_ai_search_config,
rag_store_config,
} => Tool::Retrieval {
retrieval_types,
vertex_ai_search_config,
exa_ai_search_config,
parallel_ai_search_config,
rag_store_config,
},
}),
Err(parse_error) => {
let tool_type = value
.get("type")
.and_then(|v| v.as_str())
.unwrap_or("<missing type>")
.to_string();
tracing::warn!(
"Encountered unknown Tool type '{}'. \
Parse error: {}. \
This may indicate a new API feature or a malformed response. \
The tool will be preserved in the Unknown variant.",
tool_type,
parse_error
);
Ok(Tool::Unknown {
tool_type,
data: value,
})
}
}
}
}
impl Tool {
#[must_use]
pub const fn is_unknown(&self) -> bool {
matches!(self, Self::Unknown { .. })
}
#[must_use]
pub fn unknown_tool_type(&self) -> Option<&str> {
match self {
Self::Unknown { tool_type, .. } => Some(tool_type),
_ => None,
}
}
#[must_use]
pub fn unknown_data(&self) -> Option<&serde_json::Value> {
match self {
Self::Unknown { data, .. } => Some(data),
_ => None,
}
}
}
#[derive(Clone, Serialize, Deserialize, Debug)]
pub struct FunctionDeclaration {
name: String,
description: String,
parameters: FunctionParameters,
}
#[derive(Clone, Serialize, Deserialize, Debug)]
pub struct FunctionParameters {
#[serde(rename = "type")]
type_: String,
properties: serde_json::Value,
#[serde(skip_serializing_if = "Vec::is_empty", default)]
required: Vec<String>,
}
impl FunctionDeclaration {
#[doc(hidden)]
pub fn new(name: String, description: String, parameters: FunctionParameters) -> Self {
Self {
name,
description,
parameters,
}
}
#[must_use]
pub fn builder(name: impl Into<String>) -> FunctionDeclarationBuilder {
FunctionDeclarationBuilder::new(name)
}
#[must_use]
pub fn name(&self) -> &str {
&self.name
}
#[must_use]
pub fn description(&self) -> &str {
&self.description
}
#[must_use]
pub fn parameters(&self) -> &FunctionParameters {
&self.parameters
}
#[must_use]
pub fn into_tool(self) -> Tool {
Tool::Function {
name: self.name,
description: self.description,
parameters: self.parameters,
}
}
}
impl FunctionParameters {
#[doc(hidden)]
pub fn new(type_: String, properties: serde_json::Value, required: Vec<String>) -> Self {
Self {
type_,
properties,
required,
}
}
#[must_use]
pub fn type_(&self) -> &str {
&self.type_
}
#[must_use]
pub fn properties(&self) -> &serde_json::Value {
&self.properties
}
#[must_use]
pub fn required(&self) -> &[String] {
&self.required
}
}
#[derive(Debug)]
pub struct FunctionDeclarationBuilder {
name: String,
description: String,
properties: serde_json::Value,
required: Vec<String>,
}
impl FunctionDeclarationBuilder {
pub fn new(name: impl Into<String>) -> Self {
Self {
name: name.into(),
description: String::new(),
properties: serde_json::Value::Object(serde_json::Map::new()),
required: Vec::new(),
}
}
pub fn description(mut self, description: impl Into<String>) -> Self {
self.description = description.into();
self
}
pub fn parameter(mut self, name: &str, schema: serde_json::Value) -> Self {
if let serde_json::Value::Object(ref mut map) = self.properties {
map.insert(name.to_string(), schema);
}
self
}
pub fn required(mut self, required: Vec<String>) -> Self {
self.required = required;
self
}
pub fn build(self) -> FunctionDeclaration {
if self.name.trim().is_empty() {
tracing::warn!(
"FunctionDeclaration built with empty or whitespace-only name. \
This will likely be rejected by the API."
);
}
if let serde_json::Value::Object(ref props) = self.properties {
for req in &self.required {
if !props.contains_key(req) {
tracing::warn!(
"FunctionDeclaration '{}' requires parameter '{}' which is not defined in properties. \
This will likely cause API errors.",
self.name,
req
);
}
}
}
FunctionDeclaration {
name: self.name,
description: self.description,
parameters: FunctionParameters {
type_: "object".to_string(),
properties: self.properties,
required: self.required,
},
}
}
}
#[derive(Clone, Debug, PartialEq)]
#[non_exhaustive]
pub enum FunctionCallingMode {
Auto,
Any,
None,
Validated,
Unknown {
mode_type: String,
data: serde_json::Value,
},
}
impl FunctionCallingMode {
#[must_use]
pub const fn is_unknown(&self) -> bool {
matches!(self, Self::Unknown { .. })
}
#[must_use]
pub fn unknown_mode_type(&self) -> Option<&str> {
match self {
Self::Unknown { mode_type, .. } => Some(mode_type),
_ => None,
}
}
#[must_use]
pub fn unknown_data(&self) -> Option<&serde_json::Value> {
match self {
Self::Unknown { data, .. } => Some(data),
_ => None,
}
}
}
impl Serialize for FunctionCallingMode {
fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
where
S: serde::Serializer,
{
match self {
Self::Auto => serializer.serialize_str("auto"),
Self::Any => serializer.serialize_str("any"),
Self::None => serializer.serialize_str("none"),
Self::Validated => serializer.serialize_str("validated"),
Self::Unknown { mode_type, .. } => serializer.serialize_str(mode_type),
}
}
}
impl<'de> Deserialize<'de> for FunctionCallingMode {
fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
where
D: serde::Deserializer<'de>,
{
let value = serde_json::Value::deserialize(deserializer)?;
match value.as_str() {
Some("auto") | Some("AUTO") => Ok(Self::Auto),
Some("any") | Some("ANY") => Ok(Self::Any),
Some("none") | Some("NONE") => Ok(Self::None),
Some("validated") | Some("VALIDATED") => Ok(Self::Validated),
Some(other) => {
tracing::warn!(
"Encountered unknown FunctionCallingMode '{}'. \
This may indicate a new API feature. \
The mode will be preserved in the Unknown variant.",
other
);
Ok(Self::Unknown {
mode_type: other.to_string(),
data: value,
})
}
Option::None => {
let mode_type = format!("<non-string: {}>", value);
tracing::warn!(
"FunctionCallingMode received non-string value: {}. \
Preserving in Unknown variant.",
value
);
Ok(Self::Unknown {
mode_type,
data: value,
})
}
}
}
}
#[derive(Clone, Debug, Default, PartialEq, Serialize, Deserialize)]
#[serde(default)]
pub struct AllowedTools {
#[serde(skip_serializing_if = "Option::is_none")]
pub mode: Option<FunctionCallingMode>,
#[serde(skip_serializing_if = "Vec::is_empty")]
pub tools: Vec<String>,
}
impl AllowedTools {
#[must_use]
pub fn new(tools: Vec<String>) -> Self {
Self { mode: None, tools }
}
#[must_use]
pub fn with_mode(mut self, mode: FunctionCallingMode) -> Self {
self.mode = Some(mode);
self
}
}
#[derive(Clone, Debug, PartialEq)]
#[non_exhaustive]
pub enum ToolChoice {
Mode(FunctionCallingMode),
AllowedTools(AllowedTools),
Unknown {
choice_type: String,
data: serde_json::Value,
},
}
impl ToolChoice {
#[must_use]
pub fn allowed_tools(mode: Option<FunctionCallingMode>, tools: Vec<String>) -> Self {
Self::AllowedTools(AllowedTools { mode, tools })
}
#[must_use]
pub const fn is_unknown(&self) -> bool {
matches!(self, Self::Unknown { .. })
}
#[must_use]
pub fn unknown_choice_type(&self) -> Option<&str> {
match self {
Self::Unknown { choice_type, .. } => Some(choice_type),
_ => None,
}
}
#[must_use]
pub fn unknown_data(&self) -> Option<&serde_json::Value> {
match self {
Self::Unknown { data, .. } => Some(data),
_ => None,
}
}
}
impl From<FunctionCallingMode> for ToolChoice {
fn from(mode: FunctionCallingMode) -> Self {
Self::Mode(mode)
}
}
impl Serialize for ToolChoice {
fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
where
S: serde::Serializer,
{
use serde::ser::SerializeMap;
match self {
Self::Mode(mode) => mode.serialize(serializer),
Self::AllowedTools(allowed) => {
let mut map = serializer.serialize_map(Some(1))?;
map.serialize_entry("allowed_tools", allowed)?;
map.end()
}
Self::Unknown { data, .. } => data.serialize(serializer),
}
}
}
impl<'de> Deserialize<'de> for ToolChoice {
fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
where
D: serde::Deserializer<'de>,
{
let value = serde_json::Value::deserialize(deserializer)?;
match &value {
serde_json::Value::String(_) => {
let mode: FunctionCallingMode =
serde_json::from_value(value).map_err(serde::de::Error::custom)?;
Ok(Self::Mode(mode))
}
serde_json::Value::Object(obj) if obj.contains_key("allowed_tools") => {
match serde_json::from_value::<AllowedTools>(obj["allowed_tools"].clone()) {
Ok(allowed) => Ok(Self::AllowedTools(allowed)),
Err(e) => {
tracing::warn!(
"Failed to parse tool_choice.allowed_tools: {}. \
Preserving in Unknown variant.",
e
);
Ok(Self::Unknown {
choice_type: "allowed_tools".to_string(),
data: value,
})
}
}
}
_ => {
tracing::warn!(
"Encountered unknown ToolChoice shape: {}. \
Preserving in Unknown variant.",
value
);
Ok(Self::Unknown {
choice_type: format!("<unrecognized: {}>", value),
data: value,
})
}
}
}
}
#[derive(Clone, Debug, PartialEq, Eq)]
#[non_exhaustive]
pub enum SearchType {
WebSearch,
ImageSearch,
EnterpriseWebSearch,
Unknown {
search_type: String,
data: serde_json::Value,
},
}
impl SearchType {
#[must_use]
pub const fn is_unknown(&self) -> bool {
matches!(self, Self::Unknown { .. })
}
#[must_use]
pub fn unknown_search_type(&self) -> Option<&str> {
match self {
Self::Unknown { search_type, .. } => Some(search_type),
_ => None,
}
}
#[must_use]
pub fn unknown_data(&self) -> Option<&serde_json::Value> {
match self {
Self::Unknown { data, .. } => Some(data),
_ => None,
}
}
}
impl Serialize for SearchType {
fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
where
S: serde::Serializer,
{
match self {
Self::WebSearch => serializer.serialize_str("web_search"),
Self::ImageSearch => serializer.serialize_str("image_search"),
Self::EnterpriseWebSearch => serializer.serialize_str("enterprise_web_search"),
Self::Unknown { search_type, .. } => serializer.serialize_str(search_type),
}
}
}
impl<'de> Deserialize<'de> for SearchType {
fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
where
D: serde::Deserializer<'de>,
{
let value = serde_json::Value::deserialize(deserializer)?;
match value.as_str() {
Some("web_search") => Ok(Self::WebSearch),
Some("image_search") => Ok(Self::ImageSearch),
Some("enterprise_web_search") => Ok(Self::EnterpriseWebSearch),
Some(other) => {
tracing::warn!(
"Encountered unknown SearchType '{}'. \
Preserving in Unknown variant.",
other
);
Ok(Self::Unknown {
search_type: other.to_string(),
data: value,
})
}
None => {
let search_type = format!("<non-string: {}>", value);
tracing::warn!(
"SearchType received non-string value: {}. \
Preserving in Unknown variant.",
value
);
Ok(Self::Unknown {
search_type,
data: value,
})
}
}
}
}
#[derive(Clone, Debug, PartialEq)]
#[non_exhaustive]
pub enum RetrievalType {
VertexAiSearch,
RagStore,
ExaAiSearch,
ParallelAiSearch,
Unknown {
retrieval_type: String,
data: serde_json::Value,
},
}
impl RetrievalType {
#[must_use]
pub const fn is_unknown(&self) -> bool {
matches!(self, Self::Unknown { .. })
}
#[must_use]
pub fn unknown_retrieval_type(&self) -> Option<&str> {
match self {
Self::Unknown { retrieval_type, .. } => Some(retrieval_type),
_ => None,
}
}
#[must_use]
pub fn unknown_data(&self) -> Option<&serde_json::Value> {
match self {
Self::Unknown { data, .. } => Some(data),
_ => None,
}
}
}
impl Serialize for RetrievalType {
fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
where
S: serde::Serializer,
{
match self {
Self::VertexAiSearch => serializer.serialize_str("vertex_ai_search"),
Self::RagStore => serializer.serialize_str("rag_store"),
Self::ExaAiSearch => serializer.serialize_str("exa_ai_search"),
Self::ParallelAiSearch => serializer.serialize_str("parallel_ai_search"),
Self::Unknown { retrieval_type, .. } => serializer.serialize_str(retrieval_type),
}
}
}
impl<'de> Deserialize<'de> for RetrievalType {
fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
where
D: serde::Deserializer<'de>,
{
let value = serde_json::Value::deserialize(deserializer)?;
match value.as_str() {
Some("vertex_ai_search") => Ok(Self::VertexAiSearch),
Some("rag_store") => Ok(Self::RagStore),
Some("exa_ai_search") => Ok(Self::ExaAiSearch),
Some("parallel_ai_search") => Ok(Self::ParallelAiSearch),
Some(other) => {
tracing::warn!(
"Encountered unknown RetrievalType '{}'. \
Preserving in Unknown variant.",
other
);
Ok(Self::Unknown {
retrieval_type: other.to_string(),
data: value,
})
}
None => {
let retrieval_type = format!("<non-string: {}>", value);
tracing::warn!(
"RetrievalType received non-string value: {}. \
Preserving in Unknown variant.",
value
);
Ok(Self::Unknown {
retrieval_type,
data: value,
})
}
}
}
}
#[derive(Clone, Debug, Default, PartialEq, Serialize, Deserialize)]
#[serde(default)]
pub struct VertexAiSearchConfig {
#[serde(skip_serializing_if = "Option::is_none")]
pub engine: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub datastores: Option<Vec<String>>,
}
impl VertexAiSearchConfig {
#[must_use]
pub fn new() -> Self {
Self::default()
}
#[must_use]
pub fn with_engine(mut self, engine: impl Into<String>) -> Self {
self.engine = Some(engine.into());
self
}
#[must_use]
pub fn with_datastores(mut self, datastores: Vec<String>) -> Self {
self.datastores = Some(datastores);
self
}
}
#[derive(Clone, Debug, Default, PartialEq, Serialize, Deserialize)]
#[serde(default)]
pub struct ExaAiSearchConfig {
pub api_key: String,
#[serde(skip_serializing_if = "Option::is_none")]
pub custom_config: Option<serde_json::Value>,
}
impl ExaAiSearchConfig {
#[must_use]
pub fn new(api_key: impl Into<String>) -> Self {
Self {
api_key: api_key.into(),
custom_config: None,
}
}
#[must_use]
pub fn with_custom_config(mut self, custom_config: serde_json::Value) -> Self {
self.custom_config = Some(custom_config);
self
}
}
#[derive(Clone, Debug, Default, PartialEq, Serialize, Deserialize)]
#[serde(default)]
pub struct ParallelAiSearchConfig {
#[serde(skip_serializing_if = "Option::is_none")]
pub api_key: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub custom_config: Option<serde_json::Value>,
}
impl ParallelAiSearchConfig {
#[must_use]
pub fn new() -> Self {
Self::default()
}
#[must_use]
pub fn with_api_key(mut self, api_key: impl Into<String>) -> Self {
self.api_key = Some(api_key.into());
self
}
#[must_use]
pub fn with_custom_config(mut self, custom_config: serde_json::Value) -> Self {
self.custom_config = Some(custom_config);
self
}
}
#[derive(Clone, Debug, Default, PartialEq, Serialize, Deserialize)]
#[serde(default)]
pub struct RagResource {
#[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>>,
}
impl RagResource {
#[must_use]
pub fn new(rag_corpus: impl Into<String>) -> Self {
Self {
rag_corpus: Some(rag_corpus.into()),
rag_file_ids: None,
}
}
#[must_use]
pub fn with_rag_file_ids(mut self, rag_file_ids: Vec<String>) -> Self {
self.rag_file_ids = Some(rag_file_ids);
self
}
}
#[derive(Clone, Debug, Default, PartialEq, Serialize, Deserialize)]
#[serde(default)]
pub struct HybridSearchConfig {
#[serde(skip_serializing_if = "Option::is_none")]
pub alpha: Option<f32>,
}
#[derive(Clone, Debug, Default, PartialEq, Serialize, Deserialize)]
#[serde(default)]
pub struct RagFilter {
#[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>,
#[serde(skip_serializing_if = "Option::is_none")]
pub metadata_filter: Option<String>,
}
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
pub struct RagRanking {
#[serde(default = "RagRanking::default_ranking_config")]
pub ranking_config: String,
#[serde(skip_serializing_if = "Option::is_none", default)]
pub model_name: Option<String>,
}
impl RagRanking {
fn default_ranking_config() -> String {
"rank_service".to_string()
}
#[must_use]
pub fn rank_service() -> Self {
Self::default()
}
#[must_use]
pub fn with_model_name(mut self, model_name: impl Into<String>) -> Self {
self.model_name = Some(model_name.into());
self
}
}
impl Default for RagRanking {
fn default() -> Self {
Self {
ranking_config: Self::default_ranking_config(),
model_name: None,
}
}
}
#[derive(Clone, Debug, Default, PartialEq, Serialize, Deserialize)]
#[serde(default)]
pub struct RagRetrievalConfig {
#[serde(skip_serializing_if = "Option::is_none")]
pub top_k: Option<i32>,
#[serde(skip_serializing_if = "Option::is_none")]
pub hybrid_search: Option<HybridSearchConfig>,
#[serde(rename = "filter", skip_serializing_if = "Option::is_none")]
pub filter: Option<RagFilter>,
#[serde(skip_serializing_if = "Option::is_none")]
pub ranking: Option<RagRanking>,
}
impl RagRetrievalConfig {
#[must_use]
pub fn new() -> Self {
Self::default()
}
#[must_use]
pub fn with_top_k(mut self, top_k: i32) -> Self {
self.top_k = Some(top_k);
self
}
#[must_use]
pub fn with_hybrid_search_alpha(mut self, alpha: f32) -> Self {
self.hybrid_search = Some(HybridSearchConfig { alpha: Some(alpha) });
self
}
#[must_use]
pub fn with_filter(mut self, filter: RagFilter) -> Self {
self.filter = Some(filter);
self
}
#[must_use]
pub fn with_ranking(mut self, ranking: RagRanking) -> Self {
self.ranking = Some(ranking);
self
}
}
#[derive(Clone, Debug, Default, PartialEq, Serialize, Deserialize)]
#[serde(default)]
pub struct RagStoreConfig {
#[serde(skip_serializing_if = "Option::is_none")]
pub rag_resources: Option<Vec<RagResource>>,
#[serde(skip_serializing_if = "Option::is_none")]
pub similarity_top_k: Option<i32>,
#[serde(skip_serializing_if = "Option::is_none")]
pub vector_distance_threshold: Option<f64>,
#[serde(skip_serializing_if = "Option::is_none")]
pub rag_retrieval_config: Option<RagRetrievalConfig>,
}
impl RagStoreConfig {
#[must_use]
pub fn new(rag_resources: Vec<RagResource>) -> Self {
Self {
rag_resources: Some(rag_resources),
..Default::default()
}
}
#[must_use]
pub fn with_rag_retrieval_config(mut self, config: RagRetrievalConfig) -> Self {
self.rag_retrieval_config = Some(config);
self
}
}
#[derive(Clone, Debug, Default)]
pub struct GoogleSearchConfig {
search_types: Option<Vec<SearchType>>,
}
impl GoogleSearchConfig {
#[must_use]
pub fn new() -> Self {
Self::default()
}
#[must_use]
pub fn with_search_types(mut self, search_types: Vec<SearchType>) -> Self {
self.search_types = Some(search_types);
self
}
}
impl From<GoogleSearchConfig> for Tool {
fn from(config: GoogleSearchConfig) -> Self {
Tool::GoogleSearch {
search_types: config.search_types,
}
}
}
#[derive(Clone, Debug, Default)]
pub struct GoogleMapsConfig {
enable_widget: Option<bool>,
latitude: Option<f64>,
longitude: Option<f64>,
}
impl GoogleMapsConfig {
#[must_use]
pub fn new() -> Self {
Self::default()
}
#[must_use]
pub fn with_widget(mut self) -> Self {
self.enable_widget = Some(true);
self
}
#[must_use]
pub fn with_location(mut self, latitude: f64, longitude: f64) -> Self {
self.latitude = Some(latitude);
self.longitude = Some(longitude);
self
}
}
impl From<GoogleMapsConfig> for Tool {
fn from(config: GoogleMapsConfig) -> Self {
Tool::GoogleMaps {
enable_widget: config.enable_widget,
latitude: config.latitude,
longitude: config.longitude,
}
}
}
#[derive(Clone, Debug)]
pub struct McpServerConfig {
name: String,
url: String,
allowed_tools: Option<Vec<AllowedTools>>,
headers: Option<HashMap<String, String>>,
}
impl McpServerConfig {
#[must_use]
pub fn new(name: impl Into<String>, url: impl Into<String>) -> Self {
Self {
name: name.into(),
url: url.into(),
allowed_tools: None,
headers: None,
}
}
#[must_use]
pub fn with_allowed_tools(mut self, allowed_tools: Vec<String>) -> Self {
self.allowed_tools = Some(vec![AllowedTools::new(allowed_tools)]);
self
}
#[must_use]
pub fn with_allowed_tools_config(mut self, allowed_tools: Vec<AllowedTools>) -> Self {
self.allowed_tools = Some(allowed_tools);
self
}
#[must_use]
pub fn with_headers(mut self, headers: HashMap<String, String>) -> Self {
self.headers = Some(headers);
self
}
}
impl From<McpServerConfig> for Tool {
fn from(config: McpServerConfig) -> Self {
Tool::McpServer {
name: config.name,
url: config.url,
allowed_tools: config.allowed_tools,
headers: config.headers,
}
}
}
#[derive(Clone, Debug)]
pub struct ComputerUseConfig {
environment: String,
excluded_predefined_functions: Vec<String>,
enable_prompt_injection_detection: Option<bool>,
disabled_safety_policies: Vec<String>,
}
impl ComputerUseConfig {
#[must_use]
pub fn new() -> Self {
Self {
environment: "browser".to_string(),
excluded_predefined_functions: Vec::new(),
enable_prompt_injection_detection: None,
disabled_safety_policies: Vec::new(),
}
}
#[must_use]
pub fn with_environment(mut self, environment: impl Into<String>) -> Self {
self.environment = environment.into();
self
}
#[must_use]
pub fn excluding(mut self, functions: Vec<String>) -> Self {
self.excluded_predefined_functions = functions;
self
}
#[must_use]
pub fn with_prompt_injection_detection(mut self, enabled: bool) -> Self {
self.enable_prompt_injection_detection = Some(enabled);
self
}
#[must_use]
pub fn disabling_safety_policies(mut self, policies: Vec<String>) -> Self {
self.disabled_safety_policies = policies;
self
}
}
impl Default for ComputerUseConfig {
fn default() -> Self {
Self::new()
}
}
impl From<ComputerUseConfig> for Tool {
fn from(config: ComputerUseConfig) -> Self {
Tool::ComputerUse {
environment: config.environment,
excluded_predefined_functions: config.excluded_predefined_functions,
enable_prompt_injection_detection: config.enable_prompt_injection_detection,
disabled_safety_policies: config.disabled_safety_policies,
}
}
}
#[derive(Clone, Debug)]
pub struct FileSearchConfig {
store_names: Vec<String>,
top_k: Option<i32>,
metadata_filter: Option<String>,
}
impl FileSearchConfig {
#[must_use]
pub fn new(store_names: Vec<String>) -> Self {
Self {
store_names,
top_k: None,
metadata_filter: None,
}
}
#[must_use]
pub fn with_top_k(mut self, top_k: i32) -> Self {
self.top_k = Some(top_k);
self
}
#[must_use]
pub fn with_metadata_filter(mut self, filter: impl Into<String>) -> Self {
self.metadata_filter = Some(filter.into());
self
}
}
impl From<FileSearchConfig> for Tool {
fn from(config: FileSearchConfig) -> Self {
Tool::FileSearch {
store_names: config.store_names,
top_k: config.top_k,
metadata_filter: config.metadata_filter,
}
}
}
#[derive(Clone, Debug, Default)]
pub struct RetrievalConfig {
retrieval_types: Vec<RetrievalType>,
vertex_ai_search_config: Option<VertexAiSearchConfig>,
exa_ai_search_config: Option<ExaAiSearchConfig>,
parallel_ai_search_config: Option<ParallelAiSearchConfig>,
rag_store_config: Option<Box<RagStoreConfig>>,
}
impl RetrievalConfig {
#[must_use]
pub fn new() -> Self {
Self::default()
}
fn enable(&mut self, retrieval_type: RetrievalType) {
if !self.retrieval_types.contains(&retrieval_type) {
self.retrieval_types.push(retrieval_type);
}
}
#[must_use]
pub fn with_vertex_ai_search(mut self, config: VertexAiSearchConfig) -> Self {
self.enable(RetrievalType::VertexAiSearch);
self.vertex_ai_search_config = Some(config);
self
}
#[must_use]
pub fn with_exa_ai_search(mut self, config: ExaAiSearchConfig) -> Self {
self.enable(RetrievalType::ExaAiSearch);
self.exa_ai_search_config = Some(config);
self
}
#[must_use]
pub fn with_parallel_ai_search(mut self, config: ParallelAiSearchConfig) -> Self {
self.enable(RetrievalType::ParallelAiSearch);
self.parallel_ai_search_config = Some(config);
self
}
#[must_use]
pub fn with_rag_store(mut self, config: RagStoreConfig) -> Self {
self.enable(RetrievalType::RagStore);
self.rag_store_config = Some(Box::new(config));
self
}
#[must_use]
pub fn with_retrieval_types(mut self, retrieval_types: Vec<RetrievalType>) -> Self {
self.retrieval_types = retrieval_types;
self
}
}
impl From<RetrievalConfig> for Tool {
fn from(config: RetrievalConfig) -> Self {
Tool::Retrieval {
retrieval_types: if config.retrieval_types.is_empty() {
None
} else {
Some(config.retrieval_types)
},
vertex_ai_search_config: config.vertex_ai_search_config,
exa_ai_search_config: config.exa_ai_search_config,
parallel_ai_search_config: config.parallel_ai_search_config,
rag_store_config: config.rag_store_config,
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use serde_json;
#[test]
fn test_serialize_function_declaration() {
let function = FunctionDeclaration::builder("get_weather")
.description("Get the current weather in a given location")
.parameter(
"location",
serde_json::json!({
"type": "string",
"description": "The city and state, e.g. San Francisco, CA"
}),
)
.required(vec!["location".to_string()])
.build();
let json_string = serde_json::to_string(&function).expect("Serialization failed");
let parsed: FunctionDeclaration =
serde_json::from_str(&json_string).expect("Deserialization failed");
assert_eq!(parsed.name(), "get_weather");
assert_eq!(
parsed.description(),
"Get the current weather in a given location"
);
}
#[test]
fn test_function_calling_mode_serialization() {
let test_cases = [
(FunctionCallingMode::Auto, "\"auto\""),
(FunctionCallingMode::Any, "\"any\""),
(FunctionCallingMode::None, "\"none\""),
(FunctionCallingMode::Validated, "\"validated\""),
];
for (mode, expected_json) in test_cases {
let json = serde_json::to_string(&mode).expect("Serialization failed");
assert_eq!(json, expected_json);
let parsed: FunctionCallingMode =
serde_json::from_str(&json).expect("Deserialization failed");
assert_eq!(parsed, mode);
}
for (raw, expected) in [
("\"AUTO\"", FunctionCallingMode::Auto),
("\"VALIDATED\"", FunctionCallingMode::Validated),
] {
let parsed: FunctionCallingMode =
serde_json::from_str(raw).expect("Deserialization failed");
assert_eq!(parsed, expected);
}
}
#[test]
fn test_function_calling_mode_unknown_roundtrip() {
let json = "\"FUTURE_MODE\"";
let parsed: FunctionCallingMode =
serde_json::from_str(json).expect("Deserialization failed");
assert!(parsed.is_unknown());
assert_eq!(parsed.unknown_mode_type(), Some("FUTURE_MODE"));
let reserialized = serde_json::to_string(&parsed).expect("Serialization failed");
assert_eq!(reserialized, json);
}
#[test]
fn test_function_calling_mode_helper_methods() {
assert!(!FunctionCallingMode::Auto.is_unknown());
assert!(!FunctionCallingMode::Any.is_unknown());
assert!(!FunctionCallingMode::None.is_unknown());
assert!(!FunctionCallingMode::Validated.is_unknown());
assert!(FunctionCallingMode::Auto.unknown_mode_type().is_none());
assert!(FunctionCallingMode::Auto.unknown_data().is_none());
let unknown = FunctionCallingMode::Unknown {
mode_type: "NEW_MODE".to_string(),
data: serde_json::json!("NEW_MODE"),
};
assert!(unknown.is_unknown());
assert_eq!(unknown.unknown_mode_type(), Some("NEW_MODE"));
assert!(unknown.unknown_data().is_some());
}
#[test]
fn test_function_calling_mode_non_string_value() {
let json = "123";
let parsed: FunctionCallingMode =
serde_json::from_str(json).expect("Deserialization should succeed");
assert!(parsed.is_unknown());
assert!(parsed.unknown_mode_type().unwrap().contains("<non-string:"));
}
#[test]
fn test_tool_google_search_roundtrip() {
let tool = Tool::GoogleSearch { search_types: None };
let json = serde_json::to_string(&tool).expect("Serialization failed");
assert!(json.contains("\"type\":\"google_search\""));
assert!(!json.contains("search_types"));
let parsed: Tool = serde_json::from_str(&json).expect("Deserialization failed");
assert!(matches!(parsed, Tool::GoogleSearch { .. }));
}
#[test]
fn test_tool_google_search_with_search_types_roundtrip() {
let tool = Tool::GoogleSearch {
search_types: Some(vec![SearchType::WebSearch, SearchType::ImageSearch]),
};
let json = serde_json::to_string(&tool).expect("Serialization failed");
assert!(json.contains("\"search_types\""));
assert!(json.contains("\"web_search\""));
assert!(json.contains("\"image_search\""));
let parsed: Tool = serde_json::from_str(&json).expect("Deserialization failed");
match parsed {
Tool::GoogleSearch { search_types } => {
let types = search_types.expect("Should have search_types");
assert_eq!(types.len(), 2);
assert_eq!(types[0], SearchType::WebSearch);
assert_eq!(types[1], SearchType::ImageSearch);
}
other => panic!("Expected GoogleSearch variant, got {:?}", other),
}
}
#[test]
fn test_tool_google_maps_roundtrip() {
let tool = Tool::GoogleMaps {
enable_widget: None,
latitude: None,
longitude: None,
};
let json = serde_json::to_string(&tool).expect("Serialization failed");
assert!(json.contains("\"type\":\"google_maps\""));
assert!(!json.contains("enable_widget"));
let parsed: Tool = serde_json::from_str(&json).expect("Deserialization failed");
match parsed {
Tool::GoogleMaps { enable_widget, .. } => assert_eq!(enable_widget, None),
other => panic!("Expected GoogleMaps variant, got {:?}", other),
}
}
#[test]
fn test_tool_google_maps_with_widget_roundtrip() {
let tool = Tool::GoogleMaps {
enable_widget: Some(true),
latitude: Some(40.758),
longitude: Some(-73.9855),
};
let json = serde_json::to_string(&tool).expect("Serialization failed");
assert!(json.contains("\"enable_widget\":true"));
assert!(json.contains("\"latitude\":40.758"));
assert!(json.contains("\"longitude\":-73.9855"));
let parsed: Tool = serde_json::from_str(&json).expect("Deserialization failed");
match parsed {
Tool::GoogleMaps {
enable_widget,
latitude,
longitude,
} => {
assert_eq!(enable_widget, Some(true));
assert_eq!(latitude, Some(40.758));
assert_eq!(longitude, Some(-73.9855));
}
other => panic!("Expected GoogleMaps variant, got {:?}", other),
}
}
#[test]
fn test_tool_function_roundtrip() {
let tool = Tool::Function {
name: "get_weather".to_string(),
description: "Get weather".to_string(),
parameters: FunctionParameters::new(
"object".to_string(),
serde_json::json!({}),
vec![],
),
};
let json = serde_json::to_string(&tool).expect("Serialization failed");
let parsed: Tool = serde_json::from_str(&json).expect("Deserialization failed");
match parsed {
Tool::Function { name, .. } => assert_eq!(name, "get_weather"),
other => panic!("Expected Function variant, got {:?}", other),
}
}
#[test]
fn test_tool_mcp_server_roundtrip() {
let tool = Tool::McpServer {
name: "my-server".to_string(),
url: "https://mcp.example.com/api".to_string(),
allowed_tools: None,
headers: None,
};
let json = serde_json::to_string(&tool).expect("Serialization failed");
assert!(json.contains("\"type\":\"mcp_server\""));
assert!(json.contains("\"name\":\"my-server\""));
assert!(json.contains("\"url\":\"https://mcp.example.com/api\""));
assert!(!json.contains("allowed_tools"));
assert!(!json.contains("headers"));
let parsed: Tool = serde_json::from_str(&json).expect("Deserialization failed");
match parsed {
Tool::McpServer {
name,
url,
allowed_tools,
headers,
} => {
assert_eq!(name, "my-server");
assert_eq!(url, "https://mcp.example.com/api");
assert_eq!(allowed_tools, None);
assert_eq!(headers, None);
}
other => panic!("Expected McpServer variant, got {:?}", other),
}
}
#[test]
fn test_tool_mcp_server_with_optional_fields_roundtrip() {
let tool = Tool::McpServer {
name: "my-server".to_string(),
url: "https://mcp.example.com/api".to_string(),
allowed_tools: Some(vec![
AllowedTools::new(vec!["read_file".to_string(), "list_dir".to_string()])
.with_mode(FunctionCallingMode::Auto),
]),
headers: Some(HashMap::from([(
"Authorization".to_string(),
"Bearer token".to_string(),
)])),
};
let json = serde_json::to_string(&tool).expect("Serialization failed");
assert!(json.contains("\"allowed_tools\""));
assert!(json.contains("\"headers\""));
let parsed: Tool = serde_json::from_str(&json).expect("Deserialization failed");
match parsed {
Tool::McpServer {
allowed_tools,
headers,
..
} => {
let tools = allowed_tools.expect("Should have allowed_tools");
assert_eq!(tools.len(), 1);
assert_eq!(tools[0].tools.len(), 2);
assert_eq!(tools[0].mode, Some(FunctionCallingMode::Auto));
let hdrs = headers.expect("Should have headers");
assert_eq!(hdrs.get("Authorization").unwrap(), "Bearer token");
}
other => panic!("Expected McpServer variant, got {:?}", other),
}
}
#[test]
fn test_tool_unknown_deserialization() {
let json = r#"{"type": "future_tool", "some_field": "value", "number": 42}"#;
let parsed: Tool = serde_json::from_str(json).expect("Deserialization failed");
match parsed {
Tool::Unknown { tool_type, data } => {
assert_eq!(tool_type, "future_tool");
assert_eq!(data.get("some_field").unwrap(), "value");
assert_eq!(data.get("number").unwrap(), 42);
}
_ => panic!("Expected Unknown variant"),
}
}
#[test]
fn test_tool_unknown_roundtrip() {
let tool = Tool::Unknown {
tool_type: "new_tool".to_string(),
data: serde_json::json!({"type": "new_tool", "config": {"enabled": true}}),
};
let json = serde_json::to_string(&tool).expect("Serialization failed");
assert!(json.contains("\"type\":\"new_tool\""));
assert!(json.contains("\"config\""));
let parsed: Tool = serde_json::from_str(&json).expect("Deserialization failed");
match parsed {
Tool::Unknown { tool_type, .. } => assert_eq!(tool_type, "new_tool"),
_ => panic!("Expected Unknown variant"),
}
}
#[test]
fn test_tool_unknown_helper_methods() {
let unknown_tool = Tool::Unknown {
tool_type: "future_tool".to_string(),
data: serde_json::json!({"type": "future_tool", "setting": 123}),
};
assert!(unknown_tool.is_unknown());
assert_eq!(unknown_tool.unknown_tool_type(), Some("future_tool"));
let data = unknown_tool.unknown_data().expect("Should have data");
assert_eq!(data.get("setting").unwrap(), 123);
}
#[test]
fn test_tool_computer_use_roundtrip() {
let tool = Tool::ComputerUse {
environment: "browser".to_string(),
excluded_predefined_functions: vec!["submit_form".to_string(), "download".to_string()],
enable_prompt_injection_detection: Some(true),
disabled_safety_policies: vec!["data_modification".to_string()],
};
let json = serde_json::to_string(&tool).expect("Serialization failed");
assert!(json.contains("\"type\":\"computer_use\""));
assert!(json.contains("\"environment\":\"browser\""));
assert!(json.contains("\"excluded_predefined_functions\""));
assert!(!json.contains("excludedPredefinedFunctions"));
assert!(json.contains("\"enable_prompt_injection_detection\":true"));
assert!(json.contains("\"disabled_safety_policies\""));
let parsed: Tool = serde_json::from_str(&json).expect("Deserialization failed");
match parsed {
Tool::ComputerUse {
environment,
excluded_predefined_functions,
enable_prompt_injection_detection,
disabled_safety_policies,
} => {
assert_eq!(environment, "browser");
assert_eq!(excluded_predefined_functions.len(), 2);
assert!(excluded_predefined_functions.contains(&"submit_form".to_string()));
assert_eq!(enable_prompt_injection_detection, Some(true));
assert_eq!(
disabled_safety_policies,
vec!["data_modification".to_string()]
);
}
other => panic!("Expected ComputerUse variant, got {:?}", other),
}
}
#[test]
fn test_tool_computer_use_legacy_camel_case_accepted() {
let json = r#"{"type":"computer_use","environment":"browser","excludedPredefinedFunctions":["a"]}"#;
let parsed: Tool = serde_json::from_str(json).expect("Deserialization failed");
match parsed {
Tool::ComputerUse {
excluded_predefined_functions,
..
} => assert_eq!(excluded_predefined_functions, vec!["a".to_string()]),
other => panic!("Expected ComputerUse variant, got {:?}", other),
}
}
#[test]
fn test_tool_computer_use_empty_exclusions() {
let tool = Tool::ComputerUse {
environment: "browser".to_string(),
excluded_predefined_functions: vec![],
enable_prompt_injection_detection: None,
disabled_safety_policies: vec![],
};
let json = serde_json::to_string(&tool).expect("Serialization failed");
assert!(json.contains("\"type\":\"computer_use\""));
assert!(json.contains("\"environment\":\"browser\""));
assert!(!json.contains("excluded_predefined_functions"));
let parsed: Tool = serde_json::from_str(&json).expect("Deserialization failed");
match parsed {
Tool::ComputerUse {
excluded_predefined_functions,
..
} => {
assert!(excluded_predefined_functions.is_empty());
}
other => panic!("Expected ComputerUse variant, got {:?}", other),
}
}
#[test]
fn test_tool_known_types_helper_methods() {
let google_search = Tool::GoogleSearch { search_types: None };
assert!(!google_search.is_unknown());
assert_eq!(google_search.unknown_tool_type(), None);
assert_eq!(google_search.unknown_data(), None);
let google_maps = Tool::GoogleMaps {
enable_widget: None,
latitude: None,
longitude: None,
};
assert!(!google_maps.is_unknown());
assert_eq!(google_maps.unknown_tool_type(), None);
assert_eq!(google_maps.unknown_data(), None);
let code_execution = Tool::CodeExecution;
assert!(!code_execution.is_unknown());
assert_eq!(code_execution.unknown_tool_type(), None);
assert_eq!(code_execution.unknown_data(), None);
let url_context = Tool::UrlContext;
assert!(!url_context.is_unknown());
assert_eq!(url_context.unknown_tool_type(), None);
assert_eq!(url_context.unknown_data(), None);
let computer_use = Tool::ComputerUse {
environment: "browser".to_string(),
excluded_predefined_functions: vec![],
enable_prompt_injection_detection: None,
disabled_safety_policies: vec![],
};
assert!(!computer_use.is_unknown());
assert_eq!(computer_use.unknown_tool_type(), None);
assert_eq!(computer_use.unknown_data(), None);
let function = Tool::Function {
name: "test".to_string(),
description: "Test function".to_string(),
parameters: FunctionParameters::new(
"object".to_string(),
serde_json::json!({}),
vec![],
),
};
assert!(!function.is_unknown());
assert_eq!(function.unknown_tool_type(), None);
assert_eq!(function.unknown_data(), None);
}
#[test]
fn test_tool_file_search_roundtrip() {
let tool = Tool::FileSearch {
store_names: vec!["store1".to_string(), "store2".to_string()],
top_k: Some(5),
metadata_filter: Some("category:technical".to_string()),
};
let json = serde_json::to_string(&tool).expect("Serialization failed");
assert!(json.contains("\"type\":\"file_search\""));
assert!(json.contains("\"file_search_store_names\"")); assert!(json.contains("\"top_k\":5"));
assert!(json.contains("\"metadata_filter\":\"category:technical\""));
let parsed: Tool = serde_json::from_str(&json).expect("Deserialization failed");
match parsed {
Tool::FileSearch {
store_names,
top_k,
metadata_filter,
} => {
assert_eq!(store_names, vec!["store1", "store2"]);
assert_eq!(top_k, Some(5));
assert_eq!(metadata_filter, Some("category:technical".to_string()));
}
other => panic!("Expected FileSearch variant, got {:?}", other),
}
}
#[test]
fn test_tool_file_search_minimal() {
let tool = Tool::FileSearch {
store_names: vec!["my-store".to_string()],
top_k: None,
metadata_filter: None,
};
let json = serde_json::to_string(&tool).expect("Serialization failed");
assert!(json.contains("\"type\":\"file_search\""));
assert!(json.contains("\"file_search_store_names\"")); assert!(!json.contains("\"top_k\""));
assert!(!json.contains("\"metadata_filter\""));
let parsed: Tool = serde_json::from_str(&json).expect("Deserialization failed");
match parsed {
Tool::FileSearch {
store_names,
top_k,
metadata_filter,
} => {
assert_eq!(store_names, vec!["my-store"]);
assert_eq!(top_k, None);
assert_eq!(metadata_filter, None);
}
other => panic!("Expected FileSearch variant, got {:?}", other),
}
}
#[test]
fn test_tool_file_search_helper_methods() {
let file_search = Tool::FileSearch {
store_names: vec!["store".to_string()],
top_k: None,
metadata_filter: None,
};
assert!(!file_search.is_unknown());
assert_eq!(file_search.unknown_tool_type(), None);
assert_eq!(file_search.unknown_data(), None);
}
#[test]
fn test_search_type_roundtrip() {
let types = vec![SearchType::WebSearch, SearchType::ImageSearch];
let json = serde_json::to_string(&types).expect("Serialization failed");
assert_eq!(json, r#"["web_search","image_search"]"#);
let parsed: Vec<SearchType> = serde_json::from_str(&json).expect("Deserialization failed");
assert_eq!(parsed, types);
}
#[test]
fn test_search_type_unknown_roundtrip() {
let json = r#""future_search""#;
let parsed: SearchType = serde_json::from_str(json).expect("Deserialization failed");
assert!(parsed.is_unknown());
assert_eq!(parsed.unknown_search_type(), Some("future_search"));
assert_eq!(
parsed.unknown_data(),
Some(&serde_json::Value::String("future_search".to_string()))
);
let reserialized = serde_json::to_string(&parsed).expect("Serialization failed");
assert_eq!(reserialized, json);
}
#[test]
fn test_google_search_config_into_tool() {
let tool: Tool = GoogleSearchConfig::new().into();
assert!(matches!(tool, Tool::GoogleSearch { search_types: None }));
let tool: Tool = GoogleSearchConfig::new()
.with_search_types(vec![SearchType::ImageSearch])
.into();
match tool {
Tool::GoogleSearch { search_types } => {
let types = search_types.expect("Should have search_types");
assert_eq!(types, vec![SearchType::ImageSearch]);
}
other => panic!("Expected GoogleSearch, got {:?}", other),
}
}
#[test]
fn test_google_maps_config_into_tool() {
let tool: Tool = GoogleMapsConfig::new().into();
assert!(matches!(
tool,
Tool::GoogleMaps {
enable_widget: None,
..
}
));
let tool: Tool = GoogleMapsConfig::new().with_widget().into();
assert!(matches!(
tool,
Tool::GoogleMaps {
enable_widget: Some(true),
..
}
));
}
#[test]
fn test_mcp_server_config_into_tool() {
let tool: Tool = McpServerConfig::new("server", "https://example.com").into();
match tool {
Tool::McpServer {
name,
url,
allowed_tools,
headers,
} => {
assert_eq!(name, "server");
assert_eq!(url, "https://example.com");
assert_eq!(allowed_tools, None);
assert_eq!(headers, None);
}
other => panic!("Expected McpServer, got {:?}", other),
}
}
#[test]
fn test_computer_use_config_into_tool() {
let tool: Tool = ComputerUseConfig::new().into();
match tool {
Tool::ComputerUse {
environment,
excluded_predefined_functions,
..
} => {
assert_eq!(environment, "browser");
assert!(excluded_predefined_functions.is_empty());
}
other => panic!("Expected ComputerUse, got {:?}", other),
}
let tool: Tool = ComputerUseConfig::new()
.excluding(vec!["download".to_string()])
.into();
match tool {
Tool::ComputerUse {
excluded_predefined_functions,
..
} => {
assert_eq!(excluded_predefined_functions, vec!["download"]);
}
other => panic!("Expected ComputerUse, got {:?}", other),
}
}
#[test]
fn test_file_search_config_into_tool() {
let tool: Tool = FileSearchConfig::new(vec!["store".to_string()])
.with_top_k(5)
.with_metadata_filter("cat:tech")
.into();
match tool {
Tool::FileSearch {
store_names,
top_k,
metadata_filter,
} => {
assert_eq!(store_names, vec!["store"]);
assert_eq!(top_k, Some(5));
assert_eq!(metadata_filter, Some("cat:tech".to_string()));
}
other => panic!("Expected FileSearch, got {:?}", other),
}
}
#[test]
fn test_retrieval_type_wire_roundtrip() {
for (retrieval_type, wire) in [
(RetrievalType::VertexAiSearch, "\"vertex_ai_search\""),
(RetrievalType::RagStore, "\"rag_store\""),
(RetrievalType::ExaAiSearch, "\"exa_ai_search\""),
(RetrievalType::ParallelAiSearch, "\"parallel_ai_search\""),
] {
assert_eq!(serde_json::to_string(&retrieval_type).unwrap(), wire);
let parsed: RetrievalType = serde_json::from_str(wire).unwrap();
assert_eq!(parsed, retrieval_type);
}
}
#[test]
fn test_retrieval_type_unknown_roundtrip() {
let unknown: RetrievalType = serde_json::from_str("\"bing_search\"").unwrap();
assert!(unknown.is_unknown());
assert_eq!(unknown.unknown_retrieval_type(), Some("bing_search"));
assert!(unknown.unknown_data().is_some());
assert_eq!(serde_json::to_string(&unknown).unwrap(), "\"bing_search\"");
assert!(!RetrievalType::RagStore.is_unknown());
assert_eq!(RetrievalType::RagStore.unknown_retrieval_type(), None);
assert_eq!(RetrievalType::RagStore.unknown_data(), None);
}
#[test]
fn test_tool_retrieval_vertex_ai_search_wire_shape() {
let tool: Tool = RetrievalConfig::new()
.with_vertex_ai_search(
VertexAiSearchConfig::new()
.with_engine("projects/p/locations/global/engines/e")
.with_datastores(vec!["ds-1".to_string()]),
)
.into();
let value = serde_json::to_value(&tool).unwrap();
assert_eq!(
value,
serde_json::json!({
"type": "retrieval",
"retrieval_types": ["vertex_ai_search"],
"vertex_ai_search_config": {
"engine": "projects/p/locations/global/engines/e",
"datastores": ["ds-1"]
}
})
);
}
#[test]
fn test_tool_retrieval_rag_store_wire_shape() {
let tool: Tool = RetrievalConfig::new()
.with_rag_store(
RagStoreConfig::new(vec![
RagResource::new("projects/p/locations/us/ragCorpora/c")
.with_rag_file_ids(vec!["f1".to_string()]),
])
.with_rag_retrieval_config(
RagRetrievalConfig::new()
.with_top_k(8)
.with_hybrid_search_alpha(0.5)
.with_filter(RagFilter {
vector_distance_threshold: Some(0.7),
vector_similarity_threshold: None,
metadata_filter: Some("category = \"tech\"".to_string()),
})
.with_ranking(RagRanking::rank_service().with_model_name("ranker-v2")),
),
)
.into();
let value = serde_json::to_value(&tool).unwrap();
assert_eq!(value["type"], "retrieval");
assert_eq!(value["retrieval_types"], serde_json::json!(["rag_store"]));
let rag = &value["rag_store_config"];
assert_eq!(
rag["rag_resources"][0]["rag_corpus"],
"projects/p/locations/us/ragCorpora/c"
);
assert_eq!(rag["rag_resources"][0]["rag_file_ids"][0], "f1");
let retrieval = &rag["rag_retrieval_config"];
assert_eq!(retrieval["top_k"], 8);
assert_eq!(retrieval["hybrid_search"]["alpha"], 0.5);
assert_eq!(retrieval["filter"]["vector_distance_threshold"], 0.7);
assert_eq!(
retrieval["filter"]["metadata_filter"],
"category = \"tech\""
);
assert_eq!(retrieval["ranking"]["ranking_config"], "rank_service");
assert_eq!(retrieval["ranking"]["model_name"], "ranker-v2");
}
#[test]
fn test_tool_retrieval_exa_and_parallel_wire_shape() {
let tool: Tool = RetrievalConfig::new()
.with_exa_ai_search(
ExaAiSearchConfig::new("exa-key")
.with_custom_config(serde_json::json!({"num_results": 5})),
)
.with_parallel_ai_search(ParallelAiSearchConfig::new().with_api_key("par-key"))
.into();
let value = serde_json::to_value(&tool).unwrap();
assert_eq!(
value["retrieval_types"],
serde_json::json!(["exa_ai_search", "parallel_ai_search"])
);
assert_eq!(value["exa_ai_search_config"]["api_key"], "exa-key");
assert_eq!(
value["exa_ai_search_config"]["custom_config"]["num_results"],
5
);
assert_eq!(value["parallel_ai_search_config"]["api_key"], "par-key");
}
#[test]
fn test_tool_retrieval_roundtrip() {
let tool: Tool = RetrievalConfig::new()
.with_rag_store(RagStoreConfig::new(vec![RagResource::new("corpora/c")]))
.with_vertex_ai_search(VertexAiSearchConfig::new().with_engine("engines/e"))
.into();
let json = serde_json::to_string(&tool).expect("Serialization failed");
let parsed: Tool = serde_json::from_str(&json).expect("Deserialization failed");
match parsed {
Tool::Retrieval {
retrieval_types,
vertex_ai_search_config,
rag_store_config,
exa_ai_search_config,
parallel_ai_search_config,
} => {
assert_eq!(
retrieval_types,
Some(vec![RetrievalType::RagStore, RetrievalType::VertexAiSearch])
);
assert_eq!(
vertex_ai_search_config.unwrap().engine.as_deref(),
Some("engines/e")
);
assert!(rag_store_config.is_some());
assert_eq!(exa_ai_search_config, None);
assert_eq!(parallel_ai_search_config, None);
}
other => panic!("Expected Retrieval variant, got {:?}", other),
}
}
#[test]
fn test_tool_retrieval_minimal_serializes_type_only() {
let tool = Tool::Retrieval {
retrieval_types: None,
vertex_ai_search_config: None,
exa_ai_search_config: None,
parallel_ai_search_config: None,
rag_store_config: None,
};
let json = serde_json::to_string(&tool).expect("Serialization failed");
assert_eq!(json, r#"{"type":"retrieval"}"#);
let parsed: Tool = serde_json::from_str(&json).expect("Deserialization failed");
assert!(matches!(parsed, Tool::Retrieval { .. }));
assert!(!parsed.is_unknown());
}
#[test]
fn test_retrieval_config_unknown_types_escape_hatch() {
let tool: Tool = RetrievalConfig::new()
.with_retrieval_types(vec![RetrievalType::Unknown {
retrieval_type: "future_backend".to_string(),
data: serde_json::json!("future_backend"),
}])
.into();
let value = serde_json::to_value(&tool).unwrap();
assert_eq!(value["retrieval_types"][0], "future_backend");
}
}