Implement spec gaps: reasoning tokens, tool validation, context injection, extensibility

- Anthropic adapter estimates reasoning_tokens from thinking block text lengths
- Add tool call validation against JSON schema + repair_tool_call callback
- Call adapter.initialize() on provider registration
- Extract model catalog to catalog.json data file loaded via include_str!
- Add ContentPart::Other variant for unknown/extensible content kinds
- Change StreamEvent::Error field from String to SdkError
- Add ToolContext (tool_call_id, messages, abort_signal) to execute handlers
- Gemini adapter uses gRPC status codes from error bodies for classification

Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
This commit is contained in:
Bryan Helmkamp 2026-02-20 11:59:28 -04:00
parent 5a76551847
commit 88e71c64ad
9 changed files with 684 additions and 181 deletions

View file

@ -0,0 +1,93 @@
[
{
"id": "claude-opus-4-6",
"provider": "anthropic",
"display_name": "Claude Opus 4.6",
"context_window": 200000,
"max_output": null,
"supports_tools": true,
"supports_vision": true,
"supports_reasoning": true,
"input_cost_per_million": 15.0,
"output_cost_per_million": 75.0,
"aliases": ["opus", "claude-opus"]
},
{
"id": "claude-sonnet-4-5",
"provider": "anthropic",
"display_name": "Claude Sonnet 4.5",
"context_window": 200000,
"max_output": null,
"supports_tools": true,
"supports_vision": true,
"supports_reasoning": true,
"input_cost_per_million": 3.0,
"output_cost_per_million": 15.0,
"aliases": ["sonnet", "claude-sonnet"]
},
{
"id": "gpt-5.2",
"provider": "openai",
"display_name": "GPT-5.2",
"context_window": 1047576,
"max_output": null,
"supports_tools": true,
"supports_vision": true,
"supports_reasoning": true,
"input_cost_per_million": null,
"output_cost_per_million": null,
"aliases": ["gpt5"]
},
{
"id": "gpt-5.2-mini",
"provider": "openai",
"display_name": "GPT-5.2 Mini",
"context_window": 1047576,
"max_output": null,
"supports_tools": true,
"supports_vision": true,
"supports_reasoning": true,
"input_cost_per_million": null,
"output_cost_per_million": null,
"aliases": []
},
{
"id": "gpt-5.2-codex",
"provider": "openai",
"display_name": "GPT-5.2 Codex",
"context_window": 1047576,
"max_output": null,
"supports_tools": true,
"supports_vision": true,
"supports_reasoning": true,
"input_cost_per_million": null,
"output_cost_per_million": null,
"aliases": ["codex"]
},
{
"id": "gemini-3-pro-preview",
"provider": "gemini",
"display_name": "Gemini 3 Pro (Preview)",
"context_window": 1048576,
"max_output": null,
"supports_tools": true,
"supports_vision": true,
"supports_reasoning": true,
"input_cost_per_million": null,
"output_cost_per_million": null,
"aliases": ["gemini-pro"]
},
{
"id": "gemini-3-flash-preview",
"provider": "gemini",
"display_name": "Gemini 3 Flash (Preview)",
"context_window": 1048576,
"max_output": null,
"supports_tools": true,
"supports_vision": true,
"supports_reasoning": true,
"input_cost_per_million": null,
"output_cost_per_million": null,
"aliases": ["gemini-flash"]
}
]

View file

@ -1,105 +1,11 @@
use crate::types::ModelInfo;
use std::sync::LazyLock;
/// Built-in model catalog (Section 2.9).
/// Built-in model catalog loaded from catalog.json (Section 2.9).
/// The catalog is advisory, not restrictive -- unknown model strings pass through.
static BUILT_IN_MODELS: LazyLock<Vec<ModelInfo>> = LazyLock::new(|| {
vec![
// === Anthropic ===
ModelInfo {
id: "claude-opus-4-6".into(),
provider: "anthropic".into(),
display_name: "Claude Opus 4.6".into(),
context_window: 200_000,
max_output: None,
supports_tools: true,
supports_vision: true,
supports_reasoning: true,
input_cost_per_million: Some(15.0),
output_cost_per_million: Some(75.0),
aliases: vec!["opus".into(), "claude-opus".into()],
},
ModelInfo {
id: "claude-sonnet-4-5".into(),
provider: "anthropic".into(),
display_name: "Claude Sonnet 4.5".into(),
context_window: 200_000,
max_output: None,
supports_tools: true,
supports_vision: true,
supports_reasoning: true,
input_cost_per_million: Some(3.0),
output_cost_per_million: Some(15.0),
aliases: vec!["sonnet".into(), "claude-sonnet".into()],
},
// === OpenAI ===
ModelInfo {
id: "gpt-5.2".into(),
provider: "openai".into(),
display_name: "GPT-5.2".into(),
context_window: 1_047_576,
max_output: None,
supports_tools: true,
supports_vision: true,
supports_reasoning: true,
input_cost_per_million: None,
output_cost_per_million: None,
aliases: vec!["gpt5".into()],
},
ModelInfo {
id: "gpt-5.2-mini".into(),
provider: "openai".into(),
display_name: "GPT-5.2 Mini".into(),
context_window: 1_047_576,
max_output: None,
supports_tools: true,
supports_vision: true,
supports_reasoning: true,
input_cost_per_million: None,
output_cost_per_million: None,
aliases: vec![],
},
ModelInfo {
id: "gpt-5.2-codex".into(),
provider: "openai".into(),
display_name: "GPT-5.2 Codex".into(),
context_window: 1_047_576,
max_output: None,
supports_tools: true,
supports_vision: true,
supports_reasoning: true,
input_cost_per_million: None,
output_cost_per_million: None,
aliases: vec!["codex".into()],
},
// === Gemini ===
ModelInfo {
id: "gemini-3-pro-preview".into(),
provider: "gemini".into(),
display_name: "Gemini 3 Pro (Preview)".into(),
context_window: 1_048_576,
max_output: None,
supports_tools: true,
supports_vision: true,
supports_reasoning: true,
input_cost_per_million: None,
output_cost_per_million: None,
aliases: vec!["gemini-pro".into()],
},
ModelInfo {
id: "gemini-3-flash-preview".into(),
provider: "gemini".into(),
display_name: "Gemini 3 Flash (Preview)".into(),
context_window: 1_048_576,
max_output: None,
supports_tools: true,
supports_vision: true,
supports_reasoning: true,
input_cost_per_million: None,
output_cost_per_million: None,
aliases: vec!["gemini-flash".into()],
},
]
serde_json::from_str(include_str!("catalog.json"))
.expect("embedded catalog.json must be valid")
});
/// Get model info by model ID (Section 2.9).

View file

@ -31,8 +31,11 @@ impl Client {
/// Create a Client from environment variables (Section 2.2).
/// Registers providers whose API keys are present in the environment.
/// The first registered provider becomes the default.
#[must_use]
pub fn from_env() -> Self {
///
/// # Errors
///
/// Returns `SdkError` if any provider adapter fails to initialize.
pub async fn from_env() -> Result<Self, SdkError> {
let mut client = Self {
providers: HashMap::new(),
default_provider: None,
@ -46,7 +49,7 @@ impl Client {
if let Ok(base_url) = std::env::var("ANTHROPIC_BASE_URL") {
adapter = adapter.with_base_url(base_url);
}
client.register_provider(Arc::new(adapter));
client.register_provider(Arc::new(adapter)).await?;
}
if let Ok(key) = std::env::var("OPENAI_API_KEY") {
let mut adapter = providers::OpenAiAdapter::new(key);
@ -59,7 +62,7 @@ impl Client {
if let Ok(project_id) = std::env::var("OPENAI_PROJECT_ID") {
adapter = adapter.with_project_id(project_id);
}
client.register_provider(Arc::new(adapter));
client.register_provider(Arc::new(adapter)).await?;
}
if let Ok(key) = std::env::var("GEMINI_API_KEY").or_else(|_| std::env::var("GOOGLE_API_KEY"))
{
@ -67,19 +70,28 @@ impl Client {
if let Ok(base_url) = std::env::var("GEMINI_BASE_URL") {
adapter = adapter.with_base_url(base_url);
}
client.register_provider(Arc::new(adapter));
client.register_provider(Arc::new(adapter)).await?;
}
client
Ok(client)
}
/// Register a provider adapter.
pub fn register_provider(&mut self, adapter: Arc<dyn ProviderAdapter>) {
/// Register a provider adapter. Calls `initialize()` on the adapter (Section 2.4).
///
/// # Errors
///
/// Returns `SdkError` if the adapter's `initialize()` method fails.
pub async fn register_provider(
&mut self,
adapter: Arc<dyn ProviderAdapter>,
) -> Result<(), SdkError> {
adapter.initialize().await?;
let name = adapter.name().to_string();
if self.default_provider.is_none() {
self.default_provider = Some(name.clone());
}
self.providers.insert(name, adapter);
Ok(())
}
/// Add middleware.
@ -292,7 +304,7 @@ mod tests {
#[tokio::test]
async fn complete_routes_to_default_provider() {
let mut client = Client::new(HashMap::new(), None, vec![]);
client.register_provider(Arc::new(MockProvider::new("test", "Hello!")));
client.register_provider(Arc::new(MockProvider::new("test", "Hello!"))).await.unwrap();
let response = client.complete(&test_request()).await.unwrap();
assert_eq!(response.text(), "Hello!");
@ -302,8 +314,8 @@ mod tests {
#[tokio::test]
async fn complete_routes_to_named_provider() {
let mut client = Client::new(HashMap::new(), None, vec![]);
client.register_provider(Arc::new(MockProvider::new("provider_a", "from A")));
client.register_provider(Arc::new(MockProvider::new("provider_b", "from B")));
client.register_provider(Arc::new(MockProvider::new("provider_a", "from A"))).await.unwrap();
client.register_provider(Arc::new(MockProvider::new("provider_b", "from B"))).await.unwrap();
let mut req = test_request();
req.provider = Some("provider_b".into());
@ -325,7 +337,7 @@ mod tests {
#[tokio::test]
async fn complete_errors_on_unknown_provider() {
let mut client = Client::new(HashMap::new(), None, vec![]);
client.register_provider(Arc::new(MockProvider::new("test", "Hello")));
client.register_provider(Arc::new(MockProvider::new("test", "Hello"))).await.unwrap();
let mut req = test_request();
req.provider = Some("nonexistent".into());
@ -342,10 +354,10 @@ mod tests {
let mut client = Client::new(HashMap::new(), None, vec![]);
assert_eq!(client.default_provider(), None);
client.register_provider(Arc::new(MockProvider::new("first", "1")));
client.register_provider(Arc::new(MockProvider::new("first", "1"))).await.unwrap();
assert_eq!(client.default_provider(), Some("first"));
client.register_provider(Arc::new(MockProvider::new("second", "2")));
client.register_provider(Arc::new(MockProvider::new("second", "2"))).await.unwrap();
assert_eq!(client.default_provider(), Some("first"));
}
@ -354,7 +366,7 @@ mod tests {
use futures::StreamExt;
let mut client = Client::new(HashMap::new(), None, vec![]);
client.register_provider(Arc::new(MockProvider::new("test", "streamed")));
client.register_provider(Arc::new(MockProvider::new("test", "streamed"))).await.unwrap();
let mut stream = client.stream(&test_request()).await.unwrap();
let first = stream.next().await.unwrap().unwrap();
@ -367,8 +379,8 @@ mod tests {
#[tokio::test]
async fn provider_names_returns_registered() {
let mut client = Client::new(HashMap::new(), None, vec![]);
client.register_provider(Arc::new(MockProvider::new("alpha", "")));
client.register_provider(Arc::new(MockProvider::new("beta", "")));
client.register_provider(Arc::new(MockProvider::new("alpha", ""))).await.unwrap();
client.register_provider(Arc::new(MockProvider::new("beta", ""))).await.unwrap();
let mut names = client.provider_names();
names.sort_unstable();
assert_eq!(names, vec!["alpha", "beta"]);
@ -402,7 +414,7 @@ mod tests {
#[tokio::test]
async fn middleware_wraps_complete() {
let mut client = Client::new(HashMap::new(), None, vec![]);
client.register_provider(Arc::new(MockProvider::new("test", "hello")));
client.register_provider(Arc::new(MockProvider::new("test", "hello"))).await.unwrap();
client.add_middleware(Arc::new(UppercaseMiddleware));
let response = client.complete(&test_request()).await.unwrap();

View file

@ -1,4 +1,5 @@
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
#[derive(Debug, Clone, Copy, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum ProviderErrorKind {
Authentication,
AccessDenied,
@ -27,7 +28,7 @@ impl std::fmt::Display for ProviderErrorKind {
}
}
#[derive(Debug, Clone)]
#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
pub struct ProviderErrorDetail {
pub message: String,
pub provider: String,
@ -50,7 +51,8 @@ impl ProviderErrorDetail {
}
}
#[derive(Debug, thiserror::Error)]
#[derive(Debug, Clone, serde::Serialize, serde::Deserialize, thiserror::Error)]
#[serde(tag = "type", rename_all = "snake_case")]
pub enum SdkError {
#[error("{kind} {}: {}", .detail.provider, .detail.message)]
Provider {

View file

@ -24,11 +24,13 @@ pub fn set_default_client(client: Client) {
}
/// Get the default client, lazily initialized from env.
fn get_default_client() -> Arc<Client> {
DEFAULT_CLIENT
.get()
.cloned()
.unwrap_or_else(|| Arc::new(Client::from_env()))
async fn get_default_client() -> Result<Arc<Client>, SdkError> {
if let Some(client) = DEFAULT_CLIENT.get() {
return Ok(client.clone());
}
let client = Arc::new(Client::from_env().await?);
let _ = DEFAULT_CLIENT.set(client.clone());
Ok(client)
}
fn build_initial_messages(params: &GenerateParams) -> Result<Vec<Message>, SdkError> {
@ -99,7 +101,10 @@ fn build_generate_result(steps: Vec<StepResult>, total_usage: Usage) -> Generate
/// Panics if a tool's `execute` handler is `None` when matched during tool execution.
#[allow(clippy::too_many_lines)]
pub async fn generate(params: GenerateParams) -> Result<GenerateResult, SdkError> {
let client = params.client.clone().unwrap_or_else(get_default_client);
let client = match params.client.clone() {
Some(c) => c,
None => get_default_client().await?,
};
let retry_policy = RetryPolicy {
max_retries: params.max_retries,
base_delay: 0.001,
@ -167,7 +172,7 @@ pub async fn generate(params: GenerateParams) -> Result<GenerateResult, SdkError
if tools.iter().any(|t| t.is_active()) {
let tool_refs: Vec<&Tool> =
tools.iter().map(std::convert::AsRef::as_ref).collect();
tool_results = execute_all_tools(&tool_refs, &tool_calls).await;
tool_results = execute_all_tools(&tool_refs, &tool_calls, &messages, abort_signal.as_ref()).await;
}
}
@ -570,7 +575,10 @@ pub async fn stream(params: GenerateParams) -> Result<StreamResult, SdkError> {
/// or any provider error encountered during streaming setup.
#[allow(clippy::too_many_lines)]
async fn stream_with_tool_loop(params: GenerateParams) -> Result<StreamEventStream, SdkError> {
let client = params.client.clone().unwrap_or_else(get_default_client);
let client = match params.client.clone() {
Some(c) => c,
None => get_default_client().await?,
};
let mut messages = build_initial_messages(&params)?;
let tool_definitions: Option<Vec<ToolDefinition>> = params
.tools
@ -695,7 +703,7 @@ async fn stream_with_tool_loop(params: GenerateParams) -> Result<StreamEventStre
let tool_refs: Vec<&Tool> =
tool_list.iter().map(std::convert::AsRef::as_ref).collect();
let tool_results = execute_all_tools(&tool_refs, &tool_calls).await;
let tool_results = execute_all_tools(&tool_refs, &tool_calls, &messages, abort_signal.as_ref()).await;
if tool_results.is_empty() {
return;
@ -804,7 +812,10 @@ async fn stream_generate_raw(
/// Returns `SdkError::Configuration` if both `prompt` and `messages` are set,
/// or any provider error encountered during streaming setup.
pub async fn stream_generate(params: GenerateParams) -> Result<StreamEventStream, SdkError> {
let client = params.client.clone().unwrap_or_else(get_default_client);
let client = match params.client.clone() {
Some(c) => c,
None => get_default_client().await?,
};
let messages = build_initial_messages(&params)?;
let tool_definitions: Option<Vec<ToolDefinition>> = params
.tools
@ -1188,7 +1199,7 @@ mod tests {
"get_weather",
"Get weather",
serde_json::json!({"type": "object", "properties": {"city": {"type": "string"}}}),
|args| async move {
|args, _ctx| async move {
let city = args["city"].as_str().unwrap_or("unknown");
Ok(serde_json::json!(format!("72F in {}", city)))
},
@ -1348,7 +1359,7 @@ mod tests {
"get_weather",
"Get weather",
serde_json::json!({"type": "object", "properties": {"city": {"type": "string"}}}),
|args| async move {
|args, _ctx| async move {
let city = args["city"].as_str().unwrap_or("unknown");
Ok(serde_json::json!(format!("72F in {}", city)))
},
@ -1714,7 +1725,7 @@ mod tests {
"get_weather",
"Get weather",
serde_json::json!({"type": "object", "properties": {"city": {"type": "string"}}}),
|_args| async { Ok(serde_json::json!("72F")) },
|_args, _ctx| async { Ok(serde_json::json!("72F")) },
)])
.max_tool_rounds(10)
.abort_signal(token)
@ -1776,7 +1787,7 @@ mod tests {
"get_weather",
"Get weather",
serde_json::json!({"type": "object", "properties": {"city": {"type": "string"}}}),
move |_args| {
move |_args, _ctx| {
let counter = tool_executed_clone.clone();
async move {
counter.fetch_add(1, Ordering::SeqCst);
@ -1965,7 +1976,7 @@ mod tests {
"get_weather",
"Get weather",
serde_json::json!({"type": "object", "properties": {"city": {"type": "string"}}}),
|_args| async { Ok(serde_json::json!("72F")) },
|_args, _ctx| async { Ok(serde_json::json!("72F")) },
)])
.max_tool_rounds(5)
.client(client),
@ -2019,7 +2030,7 @@ mod tests {
"get_weather",
"Get weather",
serde_json::json!({"type": "object", "properties": {"city": {"type": "string"}}}),
|_args| async { Ok(serde_json::json!("72F")) },
|_args, _ctx| async { Ok(serde_json::json!("72F")) },
)])
.max_tool_rounds(0)
.client(client),
@ -2105,7 +2116,7 @@ mod tests {
"get_weather",
"Get weather",
serde_json::json!({"type": "object", "properties": {"city": {"type": "string"}}}),
|_args| async { Ok(serde_json::json!("72F")) },
|_args, _ctx| async { Ok(serde_json::json!("72F")) },
)])
.max_tool_rounds(5)
.client(client),
@ -2168,7 +2179,7 @@ mod tests {
"get_weather",
"Get weather",
serde_json::json!({"type": "object", "properties": {"city": {"type": "string"}}}),
|_args| async { Ok(serde_json::json!("72F")) },
|_args, _ctx| async { Ok(serde_json::json!("72F")) },
)])
.max_tool_rounds(5)
.stop_when(|_steps| true) // Stop immediately after first round
@ -2295,7 +2306,7 @@ mod tests {
"get_weather",
"Get weather",
serde_json::json!({"type": "object", "properties": {"city": {"type": "string"}}}),
|_args| async { Ok(serde_json::json!("72F")) },
|_args, _ctx| async { Ok(serde_json::json!("72F")) },
)])
.max_tool_rounds(1)
.max_retries(3)
@ -2395,7 +2406,7 @@ mod tests {
"get_weather",
"Get weather",
serde_json::json!({"type": "object", "properties": {"city": {"type": "string"}}}),
|_args| async { Ok(serde_json::json!("72F")) },
|_args, _ctx| async { Ok(serde_json::json!("72F")) },
)])
.max_tool_rounds(1)
.timeout(TimeoutConfig {
@ -2528,7 +2539,7 @@ mod tests {
"get_weather",
"Get weather",
serde_json::json!({"type": "object", "properties": {"city": {"type": "string"}}}),
|_args| async { Ok(serde_json::json!("72F")) },
|_args, _ctx| async { Ok(serde_json::json!("72F")) },
)])
.max_tool_rounds(5)
.timeout(TimeoutConfig {

View file

@ -140,6 +140,24 @@ struct ApiUsage {
cache_creation_input_tokens: Option<i64>,
}
/// Estimate reasoning tokens from thinking content blocks.
/// Anthropic does not provide a separate reasoning token count,
/// so we estimate by dividing the character count of thinking text by 4.
fn estimate_reasoning_tokens(content_parts: &[ContentPart]) -> Option<i64> {
let total_chars: usize = content_parts
.iter()
.filter_map(|part| match part {
ContentPart::Thinking(td) => Some(td.text.len()),
_ => None,
})
.sum();
if total_chars > 0 {
Some((total_chars / 4).max(1) as i64)
} else {
None
}
}
fn map_finish_reason(stop_reason: Option<&str>) -> FinishReason {
match stop_reason {
Some("end_turn" | "stop_sequence") | None => FinishReason::Stop,
@ -257,6 +275,7 @@ fn content_part_to_api(part: &ContentPart) -> Option<serde_json::Value> {
ContentPart::Audio(_) => {
Some(serde_json::json!({"type": "text", "text": "[Audio content not supported by this provider]"}))
}
ContentPart::Other { .. } => None,
}
}
@ -834,6 +853,7 @@ impl StreamAccumulator {
}
fn handle_message_stop(&mut self) -> Vec<StreamEvent> {
self.usage.reasoning_tokens = estimate_reasoning_tokens(&self.content_parts);
let response = self.take_response();
vec![StreamEvent::Finish {
finish_reason: response.finish_reason.clone(),
@ -1091,6 +1111,7 @@ impl ProviderAdapter for Adapter {
map_finish_reason(api_resp.stop_reason.as_deref())
};
let total = api_resp.usage.input_tokens + api_resp.usage.output_tokens;
let reasoning_tokens = estimate_reasoning_tokens(&content_parts);
Ok(Response {
id: api_resp.id,
@ -1107,6 +1128,7 @@ impl ProviderAdapter for Adapter {
input_tokens: api_resp.usage.input_tokens,
output_tokens: api_resp.usage.output_tokens,
total_tokens: total,
reasoning_tokens,
cache_read_tokens: api_resp.usage.cache_read_input_tokens,
cache_write_tokens: api_resp.usage.cache_creation_input_tokens,
..Usage::default()

View file

@ -1,11 +1,10 @@
use base64::{engine::general_purpose::STANDARD as BASE64_STANDARD, Engine};
use futures::stream;
use crate::error::{error_from_status_code, ProviderErrorDetail, ProviderErrorKind, SdkError};
use crate::error::{error_from_grpc_status, error_from_status_code, ProviderErrorDetail, ProviderErrorKind, SdkError};
use crate::provider::{ProviderAdapter, StreamEventStream};
use crate::providers::common::{
extract_system_prompt, parse_error_body, parse_rate_limit_headers, parse_retry_after,
send_and_read_response,
};
use crate::types::{
ContentPart, FinishReason, Message, Request, Response, ResponseFormat, ResponseFormatType,
@ -462,10 +461,60 @@ fn parse_usage(metadata: Option<&UsageMetadata>) -> Usage {
})
}
/// Send an HTTP request and read the Gemini response body.
///
/// Like `send_and_read_response` but uses gRPC status code mapping when available.
async fn send_gemini_response(
request: reqwest::RequestBuilder,
) -> Result<(String, reqwest::header::HeaderMap), SdkError> {
let http_resp = request.send().await.map_err(|e| {
if e.is_timeout() {
SdkError::RequestTimeout {
message: format!("gemini: {e}"),
}
} else {
SdkError::Network {
message: e.to_string(),
}
}
})?;
let status = http_resp.status();
let retry_after = parse_retry_after(http_resp.headers());
let headers = http_resp.headers().clone();
let body = http_resp
.text()
.await
.map_err(|e| SdkError::Network {
message: e.to_string(),
})?;
if !status.is_success() {
let (msg, code, raw) = parse_error_body(&body, "status");
return Err(gemini_error(status.as_u16(), msg, code, raw, retry_after));
}
Ok((body, headers))
}
/// Map Gemini error response using gRPC status when available, falling back to HTTP status.
fn gemini_error(
status_code: u16,
msg: String,
grpc_status: Option<String>,
raw: Option<serde_json::Value>,
retry_after: Option<f64>,
) -> SdkError {
match grpc_status {
Some(grpc_code) => error_from_grpc_status(&grpc_code, msg, "gemini".to_string(), Some(grpc_code.clone()), raw, retry_after),
None => error_from_status_code(status_code, msg, "gemini".to_string(), None, raw, retry_after),
}
}
/// Send an HTTP request for streaming and return the `reqwest::Response`.
///
/// Checks for HTTP errors before returning. On error, reads the body and
/// maps it to `SdkError` using the same logic as `send_and_read_body`.
/// maps it to `SdkError` using gRPC status code mapping when available.
async fn send_streaming_request(
request: reqwest::RequestBuilder,
) -> Result<reqwest::Response, SdkError> {
@ -480,14 +529,7 @@ async fn send_streaming_request(
message: e.to_string(),
})?;
let (msg, code, raw) = parse_error_body(&body, "status");
return Err(error_from_status_code(
status.as_u16(),
msg,
"gemini".to_string(),
code,
raw,
retry_after,
));
return Err(gemini_error(status.as_u16(), msg, code, raw, retry_after));
}
Ok(http_resp)
@ -784,10 +826,8 @@ impl ProviderAdapter for Adapter {
for (key, value) in &self.default_headers {
req = req.header(key, value);
}
let (body, headers) = send_and_read_response(
let (body, headers) = send_gemini_response(
req.json(&api_body).timeout(self.request_timeout),
"gemini",
"status",
)
.await?;
@ -1062,4 +1102,38 @@ mod tests {
assert_eq!(part["inlineData"]["mimeType"], "application/pdf");
assert!(part["inlineData"]["data"].as_str().is_some());
}
#[test]
fn gemini_error_uses_grpc_status_when_available() {
use crate::error::ProviderErrorKind;
let err = gemini_error(400, "model not found".into(), Some("NOT_FOUND".into()), None, None);
assert!(matches!(err, SdkError::Provider { kind: ProviderErrorKind::NotFound, .. }));
let err = gemini_error(400, "bad args".into(), Some("INVALID_ARGUMENT".into()), None, None);
assert!(matches!(err, SdkError::Provider { kind: ProviderErrorKind::InvalidRequest, .. }));
let err = gemini_error(429, "rate limited".into(), Some("RESOURCE_EXHAUSTED".into()), None, None);
assert!(matches!(err, SdkError::Provider { kind: ProviderErrorKind::RateLimit, .. }));
let err = gemini_error(401, "bad key".into(), Some("UNAUTHENTICATED".into()), None, None);
assert!(matches!(err, SdkError::Provider { kind: ProviderErrorKind::Authentication, .. }));
let err = gemini_error(403, "denied".into(), Some("PERMISSION_DENIED".into()), None, None);
assert!(matches!(err, SdkError::Provider { kind: ProviderErrorKind::AccessDenied, .. }));
let err = gemini_error(504, "timeout".into(), Some("DEADLINE_EXCEEDED".into()), None, None);
assert!(matches!(err, SdkError::RequestTimeout { .. }));
}
#[test]
fn gemini_error_falls_back_to_http_status_without_grpc() {
use crate::error::ProviderErrorKind;
let err = gemini_error(429, "rate limited".into(), None, None, None);
assert!(matches!(err, SdkError::Provider { kind: ProviderErrorKind::RateLimit, .. }));
let err = gemini_error(500, "internal".into(), None, None, None);
assert!(matches!(err, SdkError::Provider { kind: ProviderErrorKind::Server, .. }));
}
}

View file

@ -1,11 +1,20 @@
use crate::types::{ToolCall, ToolDefinition, ToolResult};
use crate::types::{Message, ToolCall, ToolDefinition, ToolResult};
use std::future::Future;
use std::pin::Pin;
use std::sync::Arc;
use tokio_util::sync::CancellationToken;
/// Context passed to tool execute handlers (Section 5.2).
#[derive(Clone)]
pub struct ToolContext {
pub tool_call_id: String,
pub messages: Vec<Message>,
pub abort_signal: Option<CancellationToken>,
}
/// An execute handler for a tool.
pub type ExecuteHandler = Arc<
dyn Fn(serde_json::Value) -> Pin<Box<dyn Future<Output = Result<serde_json::Value, String>> + Send>>
dyn Fn(serde_json::Value, ToolContext) -> Pin<Box<dyn Future<Output = Result<serde_json::Value, String>> + Send>>
+ Send
+ Sync,
>;
@ -51,7 +60,7 @@ impl Tool {
handler: F,
) -> Self
where
F: Fn(serde_json::Value) -> Fut + Send + Sync + 'static,
F: Fn(serde_json::Value, ToolContext) -> Fut + Send + Sync + 'static,
Fut: Future<Output = Result<serde_json::Value, String>> + Send + 'static,
{
if let Err(e) = validate_tool_name(name) {
@ -63,7 +72,7 @@ impl Tool {
description: description.to_string(),
parameters,
},
execute: Some(Arc::new(move |args| Box::pin(handler(args)))),
execute: Some(Arc::new(move |args, ctx| Box::pin(handler(args, ctx)))),
}
}
@ -115,6 +124,8 @@ pub fn validate_tool_name(name: &str) -> Result<(), String> {
pub async fn execute_all_tools(
tools: &[&Tool],
tool_calls: &[ToolCall],
messages: &[Message],
abort_signal: Option<&CancellationToken>,
) -> Vec<ToolResult> {
use futures::future::join_all;
@ -125,12 +136,17 @@ pub async fn execute_all_tools(
let call_id = call.id.clone();
let call_name = call.name.clone();
let args = call.arguments.clone();
let ctx = ToolContext {
tool_call_id: call_id.clone(),
messages: messages.to_vec(),
abort_signal: abort_signal.cloned(),
};
async move {
match tool {
Some(t) if t.execute.is_some() => {
let handler = t.execute.as_ref().unwrap();
match handler(args).await {
match handler(args, ctx).await {
Ok(result) => ToolResult {
tool_call_id: call_id,
content: result,
@ -164,6 +180,162 @@ pub async fn execute_all_tools(
join_all(futures).await
}
/// A callback to repair invalid tool call arguments (Section 5.8).
/// Receives the tool call and the validation error message, returns repaired arguments
/// or an error if repair is not possible.
pub type RepairToolCallFn = Arc<
dyn Fn(
ToolCall,
String,
) -> Pin<Box<dyn Future<Output = Result<serde_json::Value, String>> + Send>>
+ Send
+ Sync,
>;
/// Validate tool call arguments against the tool's parameter schema.
/// Performs a lightweight structural check: verifies that when the schema
/// specifies `"type": "object"`, the arguments are a JSON object, and that
/// required properties are present.
fn validate_tool_args(args: &serde_json::Value, schema: &serde_json::Value) -> Result<(), String> {
let schema_type = schema.get("type").and_then(serde_json::Value::as_str);
if schema_type == Some("object") && !args.is_object() {
return Err(format!(
"Expected object arguments, got {}",
args_type_name(args)
));
}
if let (Some(obj), Some(required)) = (
args.as_object(),
schema.get("required").and_then(serde_json::Value::as_array),
) {
let missing: Vec<&str> = required
.iter()
.filter_map(serde_json::Value::as_str)
.filter(|key| !obj.contains_key(*key))
.collect();
if !missing.is_empty() {
return Err(format!("Missing required properties: {}", missing.join(", ")));
}
}
Ok(())
}
fn args_type_name(value: &serde_json::Value) -> &'static str {
match value {
serde_json::Value::Null => "null",
serde_json::Value::Bool(_) => "boolean",
serde_json::Value::Number(_) => "number",
serde_json::Value::String(_) => "string",
serde_json::Value::Array(_) => "array",
serde_json::Value::Object(_) => "object",
}
}
/// Execute all tool calls with optional schema validation and repair (Section 5.8).
///
/// Before calling a tool's execute handler, validates the arguments against the
/// tool's parameter schema. If validation fails and a `repair` callback is provided,
/// calls it to attempt repair. If repair succeeds, uses the repaired arguments.
/// If repair fails or is not configured, returns an error `ToolResult`.
pub async fn execute_all_tools_with_repair(
tools: &[&Tool],
tool_calls: &[ToolCall],
messages: &[Message],
abort_signal: Option<&CancellationToken>,
repair: Option<&RepairToolCallFn>,
) -> Vec<ToolResult> {
use futures::future::join_all;
let futures: Vec<_> = tool_calls
.iter()
.map(|call| {
let tool = tools.iter().find(|t| t.definition.name == call.name).copied();
let call_id = call.id.clone();
let call_name = call.name.clone();
let args = call.arguments.clone();
let call_clone = call.clone();
let ctx = ToolContext {
tool_call_id: call_id.clone(),
messages: messages.to_vec(),
abort_signal: abort_signal.cloned(),
};
async move {
let Some(t) = tool else {
return ToolResult {
tool_call_id: call_id,
content: serde_json::Value::String(format!("Unknown tool: {call_name}")),
is_error: true,
image_data: None,
image_media_type: None,
};
};
let Some(handler) = &t.execute else {
return ToolResult {
tool_call_id: call_id,
content: serde_json::Value::String(format!("Unknown tool: {call_name}")),
is_error: true,
image_data: None,
image_media_type: None,
};
};
let validated_args = match validate_tool_args(&args, &t.definition.parameters) {
Ok(()) => args,
Err(validation_error) => {
if let Some(repair_fn) = repair {
match repair_fn(call_clone, validation_error).await {
Ok(repaired) => repaired,
Err(repair_error) => {
return ToolResult {
tool_call_id: call_id,
content: serde_json::Value::String(format!(
"Tool call validation failed and repair failed: {repair_error}"
)),
is_error: true,
image_data: None,
image_media_type: None,
};
}
}
} else {
return ToolResult {
tool_call_id: call_id,
content: serde_json::Value::String(format!(
"Tool call validation failed: {validation_error}"
)),
is_error: true,
image_data: None,
image_media_type: None,
};
}
}
};
match handler(validated_args, ctx).await {
Ok(result) => ToolResult {
tool_call_id: call_id,
content: result,
is_error: false,
image_data: None,
image_media_type: None,
},
Err(err_msg) => ToolResult {
tool_call_id: call_id,
content: serde_json::Value::String(err_msg),
is_error: true,
image_data: None,
image_media_type: None,
},
}
}
})
.collect();
join_all(futures).await
}
#[cfg(test)]
mod tests {
use super::*;
@ -224,7 +396,7 @@ mod tests {
"test",
"test tool",
serde_json::json!({"type": "object", "properties": {}}),
|_args| async { Ok(serde_json::json!("result")) },
|_args, _ctx| async { Ok(serde_json::json!("result")) },
);
assert!(tool.is_active());
}
@ -235,7 +407,7 @@ mod tests {
"greet",
"Greet someone",
serde_json::json!({"type": "object", "properties": {"name": {"type": "string"}}}),
|args| async move {
|args, _ctx| async move {
let name = args["name"].as_str().unwrap_or("world");
Ok(serde_json::json!(format!("Hello, {}!", name)))
},
@ -248,7 +420,7 @@ mod tests {
)];
let tool_refs: Vec<&Tool> = tools.iter().collect();
let results = execute_all_tools(&tool_refs, &calls).await;
let results = execute_all_tools(&tool_refs, &calls, &[], None).await;
assert_eq!(results.len(), 1);
assert_eq!(results[0].tool_call_id, "call_1");
assert!(!results[0].is_error);
@ -266,7 +438,7 @@ mod tests {
)];
let tool_refs: Vec<&Tool> = tools.iter().collect();
let results = execute_all_tools(&tool_refs, &calls).await;
let results = execute_all_tools(&tool_refs, &calls, &[], None).await;
assert_eq!(results.len(), 1);
assert!(results[0].is_error);
assert!(results[0]
@ -282,7 +454,7 @@ mod tests {
"fail",
"Always fails",
serde_json::json!({"type": "object", "properties": {}}),
|_args| async { Err("something went wrong".to_string()) },
|_args, _ctx| async { Err("something went wrong".to_string()) },
)];
let calls = vec![ToolCall::new(
@ -292,7 +464,7 @@ mod tests {
)];
let tool_refs: Vec<&Tool> = tools.iter().collect();
let results = execute_all_tools(&tool_refs, &calls).await;
let results = execute_all_tools(&tool_refs, &calls, &[], None).await;
assert_eq!(results.len(), 1);
assert!(results[0].is_error);
assert_eq!(
@ -308,13 +480,13 @@ mod tests {
"tool_a",
"Tool A",
serde_json::json!({"type": "object", "properties": {}}),
|_args| async { Ok(serde_json::json!("result_a")) },
|_args, _ctx| async { Ok(serde_json::json!("result_a")) },
),
Tool::active(
"tool_b",
"Tool B",
serde_json::json!({"type": "object", "properties": {}}),
|_args| async { Ok(serde_json::json!("result_b")) },
|_args, _ctx| async { Ok(serde_json::json!("result_b")) },
),
];
@ -324,7 +496,7 @@ mod tests {
];
let tool_refs: Vec<&Tool> = tools.iter().collect();
let results = execute_all_tools(&tool_refs, &calls).await;
let results = execute_all_tools(&tool_refs, &calls, &[], None).await;
assert_eq!(results.len(), 2);
assert_eq!(results[0].tool_call_id, "call_1");
assert_eq!(results[0].content, serde_json::json!("result_a"));
@ -339,13 +511,13 @@ mod tests {
"succeed",
"Succeeds",
serde_json::json!({"type": "object", "properties": {}}),
|_args| async { Ok(serde_json::json!("ok")) },
|_args, _ctx| async { Ok(serde_json::json!("ok")) },
),
Tool::active(
"fail",
"Fails",
serde_json::json!({"type": "object", "properties": {}}),
|_args| async { Err("boom".to_string()) },
|_args, _ctx| async { Err("boom".to_string()) },
),
];
@ -355,7 +527,7 @@ mod tests {
];
let tool_refs: Vec<&Tool> = tools.iter().collect();
let results = execute_all_tools(&tool_refs, &calls).await;
let results = execute_all_tools(&tool_refs, &calls, &[], None).await;
assert_eq!(results.len(), 2);
assert!(!results[0].is_error);
assert!(results[1].is_error);
@ -378,7 +550,129 @@ mod tests {
"my-tool",
"bad name",
serde_json::json!({"type": "object"}),
|_args| async { Ok(serde_json::json!("result")) },
|_args, _ctx| async { Ok(serde_json::json!("result")) },
);
}
#[test]
fn validate_tool_args_valid_object() {
let schema = serde_json::json!({"type": "object", "properties": {"name": {"type": "string"}}});
let args = serde_json::json!({"name": "Alice"});
assert!(validate_tool_args(&args, &schema).is_ok());
}
#[test]
fn validate_tool_args_non_object_when_object_expected() {
let schema = serde_json::json!({"type": "object", "properties": {}});
let args = serde_json::json!("not an object");
let result = validate_tool_args(&args, &schema);
assert!(result.is_err());
assert!(result.unwrap_err().contains("Expected object"));
}
#[test]
fn validate_tool_args_missing_required_properties() {
let schema = serde_json::json!({
"type": "object",
"properties": {"name": {"type": "string"}, "age": {"type": "number"}},
"required": ["name", "age"]
});
let args = serde_json::json!({"name": "Alice"});
let result = validate_tool_args(&args, &schema);
assert!(result.is_err());
assert!(result.unwrap_err().contains("age"));
}
#[test]
fn validate_tool_args_no_schema_type_passes() {
let schema = serde_json::json!({});
let args = serde_json::json!("anything");
assert!(validate_tool_args(&args, &schema).is_ok());
}
#[tokio::test]
async fn execute_with_repair_valid_args_no_repair_needed() {
let tools = vec![Tool::active(
"greet",
"Greet someone",
serde_json::json!({"type": "object", "properties": {"name": {"type": "string"}}}),
|args, _ctx| async move {
let name = args["name"].as_str().unwrap_or("world");
Ok(serde_json::json!(format!("Hello, {}!", name)))
},
)];
let calls = vec![ToolCall::new("call_1", "greet", serde_json::json!({"name": "Alice"}))];
let tool_refs: Vec<&Tool> = tools.iter().collect();
let results = execute_all_tools_with_repair(&tool_refs, &calls, &[], None, None).await;
assert_eq!(results.len(), 1);
assert!(!results[0].is_error);
assert_eq!(results[0].content, serde_json::json!("Hello, Alice!"));
}
#[tokio::test]
async fn execute_with_repair_invalid_args_no_repair_fn() {
let tools = vec![Tool::active(
"greet",
"Greet someone",
serde_json::json!({"type": "object", "properties": {"name": {"type": "string"}}, "required": ["name"]}),
|args, _ctx| async move {
let name = args["name"].as_str().unwrap_or("world");
Ok(serde_json::json!(format!("Hello, {}!", name)))
},
)];
let calls = vec![ToolCall::new("call_1", "greet", serde_json::json!({}))];
let tool_refs: Vec<&Tool> = tools.iter().collect();
let results = execute_all_tools_with_repair(&tool_refs, &calls, &[], None, None).await;
assert_eq!(results.len(), 1);
assert!(results[0].is_error);
assert!(results[0].content.as_str().unwrap().contains("validation failed"));
}
#[tokio::test]
async fn execute_with_repair_invalid_args_repair_succeeds() {
let tools = vec![Tool::active(
"greet",
"Greet someone",
serde_json::json!({"type": "object", "properties": {"name": {"type": "string"}}, "required": ["name"]}),
|args, _ctx| async move {
let name = args["name"].as_str().unwrap_or("world");
Ok(serde_json::json!(format!("Hello, {}!", name)))
},
)];
let calls = vec![ToolCall::new("call_1", "greet", serde_json::json!({}))];
let tool_refs: Vec<&Tool> = tools.iter().collect();
let repair: RepairToolCallFn = Arc::new(|_call, _error| {
Box::pin(async { Ok(serde_json::json!({"name": "Repaired"})) })
});
let results = execute_all_tools_with_repair(&tool_refs, &calls, &[], None, Some(&repair)).await;
assert_eq!(results.len(), 1);
assert!(!results[0].is_error);
assert_eq!(results[0].content, serde_json::json!("Hello, Repaired!"));
}
#[tokio::test]
async fn execute_with_repair_invalid_args_repair_fails() {
let tools = vec![Tool::active(
"greet",
"Greet someone",
serde_json::json!({"type": "object", "properties": {"name": {"type": "string"}}, "required": ["name"]}),
|args, _ctx| async move {
let name = args["name"].as_str().unwrap_or("world");
Ok(serde_json::json!(format!("Hello, {}!", name)))
},
)];
let calls = vec![ToolCall::new("call_1", "greet", serde_json::json!({}))];
let tool_refs: Vec<&Tool> = tools.iter().collect();
let repair: RepairToolCallFn = Arc::new(|_call, _error| {
Box::pin(async { Err("cannot repair".to_string()) })
});
let results = execute_all_tools_with_repair(&tool_refs, &calls, &[], None, Some(&repair)).await;
assert_eq!(results.len(), 1);
assert!(results[0].is_error);
assert!(results[0].content.as_str().unwrap().contains("repair failed"));
}
}

View file

@ -85,8 +85,7 @@ pub struct ToolResult {
// --- 3.3 ContentPart ---
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
#[serde(tag = "kind", content = "data", rename_all = "snake_case")]
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum ContentPart {
Text(String),
Image(ImageData),
@ -96,6 +95,97 @@ pub enum ContentPart {
ToolResult(ToolResult),
Thinking(ThinkingData),
RedactedThinking(ThinkingData),
Other {
kind: String,
data: serde_json::Value,
},
}
impl Serialize for ContentPart {
fn serialize<S: serde::Serializer>(&self, serializer: S) -> Result<S::Ok, S::Error> {
use serde::ser::SerializeMap;
let mut map = serializer.serialize_map(Some(2))?;
match self {
Self::Text(v) => {
map.serialize_entry("kind", "text")?;
map.serialize_entry("data", v)?;
}
Self::Image(v) => {
map.serialize_entry("kind", "image")?;
map.serialize_entry("data", v)?;
}
Self::Audio(v) => {
map.serialize_entry("kind", "audio")?;
map.serialize_entry("data", v)?;
}
Self::Document(v) => {
map.serialize_entry("kind", "document")?;
map.serialize_entry("data", v)?;
}
Self::ToolCall(v) => {
map.serialize_entry("kind", "tool_call")?;
map.serialize_entry("data", v)?;
}
Self::ToolResult(v) => {
map.serialize_entry("kind", "tool_result")?;
map.serialize_entry("data", v)?;
}
Self::Thinking(v) => {
map.serialize_entry("kind", "thinking")?;
map.serialize_entry("data", v)?;
}
Self::RedactedThinking(v) => {
map.serialize_entry("kind", "redacted_thinking")?;
map.serialize_entry("data", v)?;
}
Self::Other { kind, data } => {
map.serialize_entry("kind", kind)?;
map.serialize_entry("data", data)?;
}
}
map.end()
}
}
impl<'de> Deserialize<'de> for ContentPart {
fn deserialize<D: serde::Deserializer<'de>>(deserializer: D) -> Result<Self, D::Error> {
let value = serde_json::Value::deserialize(deserializer)?;
let kind = value
.get("kind")
.and_then(serde_json::Value::as_str)
.ok_or_else(|| serde::de::Error::missing_field("kind"))?;
let data = value.get("data").cloned().unwrap_or(serde_json::Value::Null);
match kind {
"text" => serde_json::from_value(data)
.map(Self::Text)
.map_err(serde::de::Error::custom),
"image" => serde_json::from_value(data)
.map(Self::Image)
.map_err(serde::de::Error::custom),
"audio" => serde_json::from_value(data)
.map(Self::Audio)
.map_err(serde::de::Error::custom),
"document" => serde_json::from_value(data)
.map(Self::Document)
.map_err(serde::de::Error::custom),
"tool_call" => serde_json::from_value(data)
.map(Self::ToolCall)
.map_err(serde::de::Error::custom),
"tool_result" => serde_json::from_value(data)
.map(Self::ToolResult)
.map_err(serde::de::Error::custom),
"thinking" => serde_json::from_value(data)
.map(Self::Thinking)
.map_err(serde::de::Error::custom),
"redacted_thinking" => serde_json::from_value(data)
.map(Self::RedactedThinking)
.map_err(serde::de::Error::custom),
other => Ok(Self::Other {
kind: other.to_string(),
data,
}),
}
}
}
impl ContentPart {
@ -441,7 +531,7 @@ pub enum StreamEvent {
response: Box<Response>,
},
Error {
error: String,
error: SdkError,
raw: Option<serde_json::Value>,
},
ProviderEvent {
@ -483,11 +573,8 @@ impl StreamEvent {
}
}
pub fn error(message: impl Into<String>) -> Self {
Self::Error {
error: message.into(),
raw: None,
}
pub fn error(error: SdkError) -> Self {
Self::Error { error, raw: None }
}
}
@ -955,10 +1042,12 @@ mod tests {
#[test]
fn stream_event_error() {
let event = StreamEvent::error("something went wrong");
let event = StreamEvent::error(SdkError::Stream {
message: "something went wrong".into(),
});
match &event {
StreamEvent::Error { error, .. } => {
assert_eq!(error, "something went wrong");
assert_eq!(error.to_string(), "Stream error: something went wrong");
}
other => panic!("Expected Error, got {other:?}"),
}