mirror of
https://github.com/fabro-sh/fabro.git
synced 2026-10-09 03:20:56 +00:00
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:
parent
5a76551847
commit
88e71c64ad
9 changed files with 684 additions and 181 deletions
93
crates/unified-llm/src/catalog.json
Normal file
93
crates/unified-llm/src/catalog.json
Normal 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"]
|
||||
}
|
||||
]
|
||||
|
|
@ -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).
|
||||
|
|
|
|||
|
|
@ -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();
|
||||
|
|
|
|||
|
|
@ -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 {
|
||||
|
|
|
|||
|
|
@ -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(¶ms)?;
|
||||
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(¶ms)?;
|
||||
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 {
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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, .. }));
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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"));
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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:?}"),
|
||||
}
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue