Skip to main content

ferrin_google/
prepare_tools.rs

1//! Conversion of tool definitions and tool choice to `tools` and
2//! `toolConfig`.
3
4use ferrin_provider_util::tool_name_mapping::ToolNameMapping;
5use ferrin_spec::JsonObject;
6use ferrin_spec::JsonValue;
7use ferrin_spec::ToolChoice;
8use ferrin_spec::ToolDefinition;
9use ferrin_spec::error::ProviderError;
10use ferrin_spec::shared::Warning;
11use serde_json::json;
12
13use crate::capabilities::ModelCapabilities;
14use crate::json_schema::convert_json_schema_to_openapi_schema;
15use crate::json_schema::is_recursive_reference_error;
16
17/// Provider tool ids understood by this crate.
18pub mod ids {
19    /// Google Search grounding.
20    pub const GOOGLE_SEARCH: &str = "google.google_search";
21    /// Enterprise web search (Vertex AI).
22    pub const ENTERPRISE_WEB_SEARCH: &str = "google.enterprise_web_search";
23    /// URL context.
24    pub const URL_CONTEXT: &str = "google.url_context";
25    /// Code execution.
26    pub const CODE_EXECUTION: &str = "google.code_execution";
27    /// File search.
28    pub const FILE_SEARCH: &str = "google.file_search";
29    /// Vertex AI RAG store retrieval.
30    pub const VERTEX_RAG_STORE: &str = "google.vertex_rag_store";
31    /// Google Maps grounding.
32    pub const GOOGLE_MAPS: &str = "google.google_maps";
33}
34
35/// Wire name of the code execution tool (mapped from `google.code_execution`).
36pub const CODE_EXECUTION_TOOL_NAME: &str = "code_execution";
37
38/// `tools` and `toolConfig` of a request.
39#[derive(Debug, Clone, Default, PartialEq)]
40pub struct PreparedTools {
41    /// `tools` array.
42    pub tools: Option<Vec<JsonValue>>,
43    /// `toolConfig` object.
44    pub tool_config: Option<JsonObject>,
45    /// Warnings.
46    pub warnings: Vec<Warning>,
47}
48
49fn unsupported_tool(id: &str, details: &str) -> Warning {
50    Warning::unsupported_with_details(format!("provider tool {id}"), details)
51}
52
53const GEMINI2_ONLY: &str = "the tool is only supported with Gemini 2 and later models";
54
55fn provider_tool(
56    id: &str,
57    args: &JsonObject,
58    capabilities: ModelCapabilities,
59    warnings: &mut Vec<Warning>,
60) -> Option<JsonValue> {
61    let mut gemini2 = |wire: JsonValue| {
62        if capabilities.supports_gemini2_tools {
63            Some(wire)
64        } else {
65            warnings.push(unsupported_tool(id, GEMINI2_ONLY));
66            None
67        }
68    };
69    match id {
70        ids::GOOGLE_SEARCH => gemini2(json!({"googleSearch": JsonValue::Object(args.clone())})),
71        ids::ENTERPRISE_WEB_SEARCH => gemini2(json!({"enterpriseWebSearch": {}})),
72        ids::URL_CONTEXT => gemini2(json!({"urlContext": {}})),
73        ids::CODE_EXECUTION => gemini2(json!({"codeExecution": {}})),
74        ids::GOOGLE_MAPS => Some(json!({"googleMaps": {}})),
75        ids::FILE_SEARCH => {
76            if capabilities.supports_file_search {
77                Some(json!({"fileSearch": JsonValue::Object(args.clone())}))
78            } else {
79                warnings.push(unsupported_tool(
80                    id,
81                    "the file search tool is only supported with Gemini 2.5 and later models",
82                ));
83                None
84            }
85        }
86        ids::VERTEX_RAG_STORE => {
87            let mut store = JsonObject::new();
88            if let Some(corpus) = args.get("ragCorpus") {
89                store.insert(
90                    "rag_resources".to_owned(),
91                    json!({"rag_corpus": corpus.clone()}),
92                );
93            }
94            if let Some(top_k) = args.get("topK") {
95                store.insert("similarity_top_k".to_owned(), top_k.clone());
96            }
97            Some(json!({"retrieval": {"vertex_rag_store": store}}))
98        }
99        _ => {
100            warnings.push(Warning::unsupported(format!("provider tool {id}")));
101            None
102        }
103    }
104}
105
106fn function_declaration(
107    name: &str,
108    description: Option<&str>,
109    input_schema: &JsonValue,
110) -> Result<JsonValue, ProviderError> {
111    let mut declaration = JsonObject::new();
112    declaration.insert("name".to_owned(), JsonValue::from(name));
113    declaration.insert(
114        "description".to_owned(),
115        JsonValue::from(description.unwrap_or_default()),
116    );
117    match convert_json_schema_to_openapi_schema(input_schema) {
118        Ok(Some(parameters)) => {
119            declaration.insert("parameters".to_owned(), parameters);
120        }
121        Ok(None) => {}
122        Err(error) if is_recursive_reference_error(&error) => {
123            declaration.insert("parametersJsonSchema".to_owned(), input_schema.clone());
124        }
125        Err(error) => return Err(error.into()),
126    }
127    Ok(JsonValue::Object(declaration))
128}
129
130fn function_calling_config(mode: &str, allowed: Option<&str>) -> JsonObject {
131    let mut config = JsonObject::new();
132    config.insert("mode".to_owned(), JsonValue::from(mode));
133    if let Some(name) = allowed {
134        config.insert("allowedFunctionNames".to_owned(), json!([name]));
135    }
136    config
137}
138
139/// Converts `tools` and `tool_choice`.
140///
141/// # Errors
142///
143/// Returns [`ProviderError::UnsupportedFunctionality`] when a function tool
144/// schema cannot be converted (other than recursive references, which fall
145/// back to `parametersJsonSchema`).
146pub fn prepare_tools(
147    tools: &[ToolDefinition],
148    tool_choice: Option<&ToolChoice>,
149    capabilities: ModelCapabilities,
150    mapping: &ToolNameMapping,
151    retrieval_config: Option<&JsonObject>,
152) -> Result<PreparedTools, ProviderError> {
153    let mut prepared = PreparedTools::default();
154    if tools.is_empty() {
155        prepared.tool_config = retrieval_config.map(|config| {
156            let mut tool_config = JsonObject::new();
157            tool_config.insert(
158                "retrievalConfig".to_owned(),
159                JsonValue::Object(config.clone()),
160            );
161            tool_config
162        });
163        return Ok(prepared);
164    }
165    let mut declarations = Vec::new();
166    let mut provider_tools = Vec::new();
167    let mut any_strict = false;
168    for tool in tools {
169        match tool {
170            ToolDefinition::Function {
171                name,
172                description,
173                input_schema,
174                strict,
175                ..
176            } => {
177                any_strict |= *strict == Some(true);
178                declarations.push(function_declaration(
179                    mapping.to_provider_tool_name(name.as_str()),
180                    description.as_deref(),
181                    input_schema,
182                )?);
183            }
184            ToolDefinition::Provider { id, args, .. } => {
185                if let Some(wire) = provider_tool(id, args, capabilities, &mut prepared.warnings) {
186                    provider_tools.push(wire);
187                }
188            }
189            #[allow(unreachable_patterns, reason = "ToolDefinition is non-exhaustive")]
190            _ => prepared
191                .warnings
192                .push(Warning::unsupported("tool definition")),
193        }
194    }
195    let has_functions = !declarations.is_empty();
196    let has_provider_tools = !provider_tools.is_empty();
197    let mut tool_config: Option<JsonObject> = None;
198    if has_functions && has_provider_tools {
199        if capabilities.uses_gemini3_features {
200            let mut wire = provider_tools;
201            wire.push(json!({"functionDeclarations": declarations}));
202            prepared.tools = Some(wire);
203            let calling = match tool_choice {
204                Some(ToolChoice::None) => function_calling_config("NONE", None),
205                Some(ToolChoice::Required) => function_calling_config("ANY", None),
206                Some(ToolChoice::Tool { tool_name }) => function_calling_config(
207                    "ANY",
208                    Some(mapping.to_provider_tool_name(tool_name.as_str())),
209                ),
210                _ => function_calling_config("VALIDATED", None),
211            };
212            let mut config = JsonObject::new();
213            config.insert(
214                "functionCallingConfig".to_owned(),
215                JsonValue::Object(calling),
216            );
217            config.insert(
218                "includeServerSideToolInvocations".to_owned(),
219                JsonValue::Bool(true),
220            );
221            tool_config = Some(config);
222        } else {
223            prepared.warnings.push(Warning::unsupported_with_details(
224                "combination of function and provider-defined tools",
225                "function tools and provider-defined tools cannot be mixed in one request with this model; only the provider-defined tools were sent",
226            ));
227            prepared.tools = Some(provider_tools);
228        }
229    } else if has_provider_tools {
230        prepared.tools = Some(provider_tools);
231    } else if has_functions {
232        prepared.tools = Some(vec![json!({"functionDeclarations": declarations})]);
233        let strict_mode = if any_strict { "VALIDATED" } else { "AUTO" };
234        let calling = match tool_choice {
235            None => any_strict.then(|| function_calling_config("VALIDATED", None)),
236            Some(ToolChoice::Auto) => Some(function_calling_config(strict_mode, None)),
237            Some(ToolChoice::None) => Some(function_calling_config("NONE", None)),
238            Some(ToolChoice::Required) => Some(function_calling_config("ANY", None)),
239            Some(ToolChoice::Tool { tool_name }) => Some(function_calling_config(
240                "ANY",
241                Some(mapping.to_provider_tool_name(tool_name.as_str())),
242            )),
243            #[allow(unreachable_patterns, reason = "ToolChoice is non-exhaustive")]
244            Some(_) => None,
245        };
246        if let Some(calling) = calling {
247            let mut config = JsonObject::new();
248            config.insert(
249                "functionCallingConfig".to_owned(),
250                JsonValue::Object(calling),
251            );
252            tool_config = Some(config);
253        }
254    }
255    if let Some(retrieval) = retrieval_config {
256        tool_config.get_or_insert_with(JsonObject::new).insert(
257            "retrievalConfig".to_owned(),
258            JsonValue::Object(retrieval.clone()),
259        );
260    }
261    prepared.tool_config = tool_config;
262    Ok(prepared)
263}