mirror of
https://github.com/fabro-sh/fabro.git
synced 2026-10-10 03:30:59 +00:00
Implement spec gaps: rate limit headers, error classification, total timeout, metadata, stream_object
- Parse x-ratelimit-* headers into RateLimitInfo for Anthropic, OpenAI, and OpenAI-compatible providers (previously hardcoded to None) - Add "not found"/"does not exist" and "unauthorized"/"invalid key" error message classification patterns for ambiguous HTTP status codes - Apply TimeoutConfig.total to wrap the entire multi-step generate() loop (previously only per_step was used) - Add metadata field to GenerateParams with builder method, pass through to Request instead of hardcoding None - Implement stream_object() for streaming structured output with incremental JSON parsing via new ObjectStreamEvent type (Partial/Delta/Complete variants) - Add OpenAI-compatible Chat Completions adapter for third-party endpoints Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
This commit is contained in:
parent
cfd1d5ccea
commit
ed07d43335
16 changed files with 4880 additions and 294 deletions
19
Cargo.lock
generated
19
Cargo.lock
generated
|
|
@ -880,6 +880,7 @@ dependencies = [
|
|||
"bytes",
|
||||
"encoding_rs",
|
||||
"futures-core",
|
||||
"futures-util",
|
||||
"h2",
|
||||
"http",
|
||||
"http-body",
|
||||
|
|
@ -901,12 +902,14 @@ dependencies = [
|
|||
"sync_wrapper",
|
||||
"tokio",
|
||||
"tokio-native-tls",
|
||||
"tokio-util",
|
||||
"tower",
|
||||
"tower-http",
|
||||
"tower-service",
|
||||
"url",
|
||||
"wasm-bindgen",
|
||||
"wasm-bindgen-futures",
|
||||
"wasm-streams",
|
||||
"web-sys",
|
||||
]
|
||||
|
||||
|
|
@ -1386,8 +1389,11 @@ version = "0.1.0"
|
|||
dependencies = [
|
||||
"anyhow",
|
||||
"async-trait",
|
||||
"base64",
|
||||
"bytes",
|
||||
"dotenvy",
|
||||
"futures",
|
||||
"http",
|
||||
"rand",
|
||||
"reqwest",
|
||||
"serde",
|
||||
|
|
@ -1553,6 +1559,19 @@ dependencies = [
|
|||
"wasmparser",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "wasm-streams"
|
||||
version = "0.4.2"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "15053d8d85c7eccdbefef60f06769760a563c7f0a9d6902a13d35c7800b0ad65"
|
||||
dependencies = [
|
||||
"futures-util",
|
||||
"js-sys",
|
||||
"wasm-bindgen",
|
||||
"wasm-bindgen-futures",
|
||||
"web-sys",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "wasmparser"
|
||||
version = "0.244.0"
|
||||
|
|
|
|||
|
|
@ -20,10 +20,12 @@ thiserror = "2"
|
|||
serde = { version = "1", features = ["derive"] }
|
||||
serde_json = "1"
|
||||
tokio = { version = "1", features = ["full"] }
|
||||
reqwest = { version = "0.12", features = ["json"] }
|
||||
reqwest = { version = "0.12", features = ["json", "stream"] }
|
||||
uuid = { version = "1", features = ["v4"] }
|
||||
rand = "0.8"
|
||||
dotenvy = "0.15"
|
||||
futures = "0.3"
|
||||
tokio-stream = "0.1"
|
||||
async-trait = "0.1"
|
||||
base64 = "0.22"
|
||||
bytes = "1"
|
||||
|
|
|
|||
|
|
@ -21,9 +21,12 @@ futures.workspace = true
|
|||
tokio-stream.workspace = true
|
||||
async-trait.workspace = true
|
||||
reqwest.workspace = true
|
||||
base64.workspace = true
|
||||
bytes.workspace = true
|
||||
|
||||
[dev-dependencies]
|
||||
dotenvy.workspace = true
|
||||
http = "1"
|
||||
tokio = { workspace = true, features = ["test-util", "macros"] }
|
||||
|
||||
[lints]
|
||||
|
|
|
|||
|
|
@ -1,6 +1,7 @@
|
|||
use crate::error::SdkError;
|
||||
use crate::middleware::{Middleware, NextFn, NextStreamFn};
|
||||
use crate::provider::{ProviderAdapter, StreamEventStream};
|
||||
use crate::providers;
|
||||
use crate::types::{Request, Response};
|
||||
use std::collections::HashMap;
|
||||
use std::sync::Arc;
|
||||
|
|
@ -32,13 +33,38 @@ impl Client {
|
|||
/// The first registered provider becomes the default.
|
||||
#[must_use]
|
||||
pub fn from_env() -> Self {
|
||||
// In a real implementation, this would check for OPENAI_API_KEY, ANTHROPIC_API_KEY, etc.
|
||||
// and register the appropriate adapters. For now, return an empty client.
|
||||
Self {
|
||||
let mut client = Self {
|
||||
providers: HashMap::new(),
|
||||
default_provider: None,
|
||||
middleware: Vec::new(),
|
||||
};
|
||||
|
||||
// Register providers whose API keys are present in the environment.
|
||||
// Order determines which becomes the default provider.
|
||||
if let Ok(key) = std::env::var("ANTHROPIC_API_KEY") {
|
||||
let mut adapter = providers::AnthropicAdapter::new(key);
|
||||
if let Ok(base_url) = std::env::var("ANTHROPIC_BASE_URL") {
|
||||
adapter = adapter.with_base_url(base_url);
|
||||
}
|
||||
client.register_provider(Arc::new(adapter));
|
||||
}
|
||||
if let Ok(key) = std::env::var("OPENAI_API_KEY") {
|
||||
let mut adapter = providers::OpenAiAdapter::new(key);
|
||||
if let Ok(base_url) = std::env::var("OPENAI_BASE_URL") {
|
||||
adapter = adapter.with_base_url(base_url);
|
||||
}
|
||||
client.register_provider(Arc::new(adapter));
|
||||
}
|
||||
if let Ok(key) = std::env::var("GEMINI_API_KEY").or_else(|_| std::env::var("GOOGLE_API_KEY"))
|
||||
{
|
||||
let mut adapter = providers::GeminiAdapter::new(key);
|
||||
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 a provider adapter.
|
||||
|
|
|
|||
|
|
@ -140,6 +140,18 @@ pub fn error_from_status_code(
|
|||
|
||||
// First check message-based classification for ambiguous cases
|
||||
let lower_msg = detail.message.to_lowercase();
|
||||
if lower_msg.contains("not found") || lower_msg.contains("does not exist") {
|
||||
return SdkError::Provider {
|
||||
kind: ProviderErrorKind::NotFound,
|
||||
detail: Box::new(detail),
|
||||
};
|
||||
}
|
||||
if lower_msg.contains("unauthorized") || lower_msg.contains("invalid key") {
|
||||
return SdkError::Provider {
|
||||
kind: ProviderErrorKind::Authentication,
|
||||
detail: Box::new(detail),
|
||||
};
|
||||
}
|
||||
if lower_msg.contains("context length") || lower_msg.contains("too many tokens") {
|
||||
return SdkError::Provider {
|
||||
kind: ProviderErrorKind::ContextLength,
|
||||
|
|
@ -372,6 +384,58 @@ mod tests {
|
|||
assert!(matches!(err, SdkError::Provider { kind: ProviderErrorKind::ContentFilter, .. }));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn error_message_classification_not_found() {
|
||||
let err = error_from_status_code(
|
||||
400,
|
||||
"The model gpt-5 was not found".into(),
|
||||
"openai".into(),
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
);
|
||||
assert!(matches!(err, SdkError::Provider { kind: ProviderErrorKind::NotFound, .. }));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn error_message_classification_does_not_exist() {
|
||||
let err = error_from_status_code(
|
||||
400,
|
||||
"The resource does not exist".into(),
|
||||
"openai".into(),
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
);
|
||||
assert!(matches!(err, SdkError::Provider { kind: ProviderErrorKind::NotFound, .. }));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn error_message_classification_unauthorized() {
|
||||
let err = error_from_status_code(
|
||||
400,
|
||||
"Request unauthorized for this resource".into(),
|
||||
"openai".into(),
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
);
|
||||
assert!(matches!(err, SdkError::Provider { kind: ProviderErrorKind::Authentication, .. }));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn error_message_classification_invalid_key() {
|
||||
let err = error_from_status_code(
|
||||
400,
|
||||
"Provided invalid key for authentication".into(),
|
||||
"openai".into(),
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
);
|
||||
assert!(matches!(err, SdkError::Provider { kind: ProviderErrorKind::Authentication, .. }));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn grpc_status_mapping() {
|
||||
let err = error_from_grpc_status("NOT_FOUND", "model not found".into(), "gemini".into(), None, None, None);
|
||||
|
|
|
|||
|
|
@ -4,9 +4,12 @@ use crate::provider::StreamEventStream;
|
|||
use crate::retry::retry;
|
||||
use crate::tools::{execute_all_tools, Tool};
|
||||
use crate::types::{
|
||||
FinishReason, GenerateResult, Message, Request, Response, ResponseFormat, ResponseFormatType,
|
||||
RetryPolicy, StepResult, StreamEvent, ToolCall, ToolChoice, ToolDefinition, Usage,
|
||||
FinishReason, GenerateResult, Message, ObjectStreamEvent, Request, Response, ResponseFormat,
|
||||
ResponseFormatType, RetryPolicy, StepResult, StreamEvent, TimeoutConfig, ToolCall, ToolChoice,
|
||||
ToolDefinition, Usage,
|
||||
};
|
||||
use futures::StreamExt;
|
||||
use std::pin::Pin;
|
||||
use std::sync::Arc;
|
||||
use tokio::sync::OnceCell;
|
||||
|
||||
|
|
@ -61,7 +64,7 @@ fn build_request(
|
|||
max_tokens: params.max_tokens,
|
||||
stop_sequences: params.stop_sequences.clone(),
|
||||
reasoning_effort: params.reasoning_effort.clone(),
|
||||
metadata: None,
|
||||
metadata: params.metadata.clone(),
|
||||
provider_options: params.provider_options.clone(),
|
||||
}
|
||||
}
|
||||
|
|
@ -108,69 +111,101 @@ pub async fn generate(params: GenerateParams) -> Result<GenerateResult, SdkError
|
|||
.map(|tools| tools.iter().map(|t| t.definition.clone()).collect());
|
||||
|
||||
let max_tool_rounds = params.max_tool_rounds;
|
||||
let mut steps: Vec<StepResult> = Vec::new();
|
||||
let mut total_usage = Usage::default();
|
||||
|
||||
let mut round = 0u32;
|
||||
loop {
|
||||
let request = build_request(¶ms, &messages, tool_definitions.as_deref());
|
||||
let generate_future = async {
|
||||
let mut steps: Vec<StepResult> = Vec::new();
|
||||
let mut total_usage = Usage::default();
|
||||
|
||||
let client_ref = client.clone();
|
||||
let response = retry(&retry_policy, || {
|
||||
let c = client_ref.clone();
|
||||
let r = request.clone();
|
||||
async move { c.complete(&r).await }
|
||||
})
|
||||
.await?;
|
||||
let mut round = 0u32;
|
||||
loop {
|
||||
let request = build_request(¶ms, &messages, tool_definitions.as_deref());
|
||||
|
||||
let tool_calls = response.tool_calls();
|
||||
let mut tool_results = Vec::new();
|
||||
let client_ref = client.clone();
|
||||
let response =
|
||||
if let Some(per_step) = params.timeout.as_ref().and_then(|t| t.per_step) {
|
||||
let duration = std::time::Duration::from_secs_f64(per_step);
|
||||
tokio::time::timeout(duration, retry(&retry_policy, || {
|
||||
let c = client_ref.clone();
|
||||
let r = request.clone();
|
||||
async move { c.complete(&r).await }
|
||||
}))
|
||||
.await
|
||||
.map_err(|_| SdkError::RequestTimeout {
|
||||
message: format!("Per-step timeout of {per_step}s exceeded"),
|
||||
})?
|
||||
} else {
|
||||
retry(&retry_policy, || {
|
||||
let c = client_ref.clone();
|
||||
let r = request.clone();
|
||||
async move { c.complete(&r).await }
|
||||
})
|
||||
.await
|
||||
}?;
|
||||
|
||||
if !tool_calls.is_empty()
|
||||
&& response.finish_reason == FinishReason::ToolCalls
|
||||
&& params.tools.is_some()
|
||||
{
|
||||
let tools = params.tools.as_ref().expect("checked above");
|
||||
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;
|
||||
let tool_calls = response.tool_calls();
|
||||
let mut tool_results = Vec::new();
|
||||
|
||||
if !tool_calls.is_empty()
|
||||
&& response.finish_reason == FinishReason::ToolCalls
|
||||
&& params.tools.is_some()
|
||||
{
|
||||
let tools = params.tools.as_ref().expect("checked above");
|
||||
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;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
total_usage = total_usage + response.usage.clone();
|
||||
total_usage = total_usage + response.usage.clone();
|
||||
|
||||
let should_continue = !tool_calls.is_empty()
|
||||
&& response.finish_reason == FinishReason::ToolCalls
|
||||
&& round < max_tool_rounds
|
||||
&& !tool_results.is_empty();
|
||||
steps.push(StepResult {
|
||||
response,
|
||||
tool_results,
|
||||
});
|
||||
|
||||
if should_continue {
|
||||
messages.push(response.message.clone());
|
||||
for result in &tool_results {
|
||||
let last = steps.last().expect("just pushed");
|
||||
let should_continue = !tool_calls.is_empty()
|
||||
&& last.response.finish_reason == FinishReason::ToolCalls
|
||||
&& round < max_tool_rounds
|
||||
&& !last.tool_results.is_empty()
|
||||
&& !params.stop_when.as_ref().is_some_and(|f| f(&steps));
|
||||
|
||||
if !should_continue {
|
||||
break;
|
||||
}
|
||||
|
||||
let last = steps.last().expect("just pushed");
|
||||
messages.push(last.response.message.clone());
|
||||
for result in &last.tool_results {
|
||||
messages.push(Message::tool_result(
|
||||
&result.tool_call_id,
|
||||
result.content.to_string(),
|
||||
result.is_error,
|
||||
));
|
||||
}
|
||||
|
||||
round += 1;
|
||||
}
|
||||
|
||||
steps.push(StepResult {
|
||||
response,
|
||||
tool_results,
|
||||
});
|
||||
Ok(build_generate_result(steps, total_usage))
|
||||
};
|
||||
|
||||
if !should_continue {
|
||||
break;
|
||||
}
|
||||
|
||||
round += 1;
|
||||
if let Some(total) = params.timeout.as_ref().and_then(|t| t.total) {
|
||||
let duration = std::time::Duration::from_secs_f64(total);
|
||||
tokio::time::timeout(duration, generate_future)
|
||||
.await
|
||||
.map_err(|_| SdkError::RequestTimeout {
|
||||
message: format!("Total timeout of {total}s exceeded"),
|
||||
})?
|
||||
} else {
|
||||
generate_future.await
|
||||
}
|
||||
|
||||
Ok(build_generate_result(steps, total_usage))
|
||||
}
|
||||
|
||||
/// Callback type for custom stop conditions in the tool loop.
|
||||
pub type StopCondition = Arc<dyn Fn(&[StepResult]) -> bool + Send + Sync>;
|
||||
|
||||
/// Parameters for `generate()` (Section 4.3).
|
||||
#[derive(Clone)]
|
||||
pub struct GenerateParams {
|
||||
|
|
@ -189,8 +224,12 @@ pub struct GenerateParams {
|
|||
pub reasoning_effort: Option<String>,
|
||||
pub provider: Option<String>,
|
||||
pub provider_options: Option<serde_json::Value>,
|
||||
pub metadata: Option<std::collections::HashMap<String, String>>,
|
||||
pub max_retries: u32,
|
||||
pub timeout: Option<TimeoutConfig>,
|
||||
pub client: Option<Arc<Client>>,
|
||||
/// Custom stop condition checked after each tool round (Section 4.3).
|
||||
pub stop_when: Option<StopCondition>,
|
||||
}
|
||||
|
||||
impl GenerateParams {
|
||||
|
|
@ -211,8 +250,11 @@ impl GenerateParams {
|
|||
reasoning_effort: None,
|
||||
provider: None,
|
||||
provider_options: None,
|
||||
metadata: None,
|
||||
max_retries: 2,
|
||||
timeout: None,
|
||||
client: None,
|
||||
stop_when: None,
|
||||
}
|
||||
}
|
||||
|
||||
|
|
@ -257,6 +299,85 @@ impl GenerateParams {
|
|||
self.provider = Some(provider.into());
|
||||
self
|
||||
}
|
||||
|
||||
#[must_use]
|
||||
pub fn tool_choice(mut self, tool_choice: ToolChoice) -> Self {
|
||||
self.tool_choice = Some(tool_choice);
|
||||
self
|
||||
}
|
||||
|
||||
#[must_use]
|
||||
pub fn response_format(mut self, response_format: ResponseFormat) -> Self {
|
||||
self.response_format = Some(response_format);
|
||||
self
|
||||
}
|
||||
|
||||
#[must_use]
|
||||
pub const fn temperature(mut self, temperature: f64) -> Self {
|
||||
self.temperature = Some(temperature);
|
||||
self
|
||||
}
|
||||
|
||||
#[must_use]
|
||||
pub const fn top_p(mut self, top_p: f64) -> Self {
|
||||
self.top_p = Some(top_p);
|
||||
self
|
||||
}
|
||||
|
||||
#[must_use]
|
||||
pub const fn max_tokens(mut self, max_tokens: i64) -> Self {
|
||||
self.max_tokens = Some(max_tokens);
|
||||
self
|
||||
}
|
||||
|
||||
#[must_use]
|
||||
pub fn stop_sequences(mut self, stop_sequences: Vec<String>) -> Self {
|
||||
self.stop_sequences = Some(stop_sequences);
|
||||
self
|
||||
}
|
||||
|
||||
#[must_use]
|
||||
pub fn reasoning_effort(mut self, reasoning_effort: impl Into<String>) -> Self {
|
||||
self.reasoning_effort = Some(reasoning_effort.into());
|
||||
self
|
||||
}
|
||||
|
||||
#[must_use]
|
||||
pub fn provider_options(mut self, provider_options: serde_json::Value) -> Self {
|
||||
self.provider_options = Some(provider_options);
|
||||
self
|
||||
}
|
||||
|
||||
#[must_use]
|
||||
pub fn metadata(mut self, metadata: std::collections::HashMap<String, String>) -> Self {
|
||||
self.metadata = Some(metadata);
|
||||
self
|
||||
}
|
||||
|
||||
#[must_use]
|
||||
pub const fn max_retries(mut self, max_retries: u32) -> Self {
|
||||
self.max_retries = max_retries;
|
||||
self
|
||||
}
|
||||
|
||||
#[must_use]
|
||||
pub const fn timeout(mut self, timeout: TimeoutConfig) -> Self {
|
||||
self.timeout = Some(timeout);
|
||||
self
|
||||
}
|
||||
|
||||
/// Set a custom stop condition for the tool loop (Section 4.3).
|
||||
///
|
||||
/// The callback receives the accumulated steps so far and returns `true`
|
||||
/// to stop the tool loop early.
|
||||
#[must_use]
|
||||
pub fn stop_when(
|
||||
mut self,
|
||||
f: impl Fn(&[StepResult]) -> bool + Send + Sync + 'static,
|
||||
) -> Self {
|
||||
self.stop_when = Some(Arc::new(f));
|
||||
self
|
||||
}
|
||||
}
|
||||
|
||||
/// `StreamAccumulator` collects stream events into a complete Response (Section 4.4).
|
||||
|
|
@ -339,6 +460,19 @@ impl Default for StreamAccumulator {
|
|||
///
|
||||
/// Returns `SdkError::Configuration` if both `prompt` and `messages` are set,
|
||||
/// or any provider error encountered during streaming setup.
|
||||
pub async fn stream(params: GenerateParams) -> Result<StreamEventStream, SdkError> {
|
||||
stream_generate(params).await
|
||||
}
|
||||
|
||||
/// High-level streaming generation (Section 4.4).
|
||||
/// Returns a `StreamEventStream` that the caller can iterate over.
|
||||
///
|
||||
/// Alias: prefer [`stream()`] for consistency with the spec.
|
||||
///
|
||||
/// # Errors
|
||||
///
|
||||
/// 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 messages = build_initial_messages(¶ms)?;
|
||||
|
|
@ -384,6 +518,94 @@ pub async fn generate_object(
|
|||
}
|
||||
}
|
||||
|
||||
/// Stream type for `stream_object()`.
|
||||
pub type ObjectStream =
|
||||
Pin<Box<dyn futures::Stream<Item = Result<ObjectStreamEvent, SdkError>> + Send>>;
|
||||
|
||||
/// Streaming structured output with incremental JSON parsing (Section 4.6).
|
||||
///
|
||||
/// Combines streaming with structured output: sets `response_format` to `json_schema`,
|
||||
/// streams the response, and attempts to parse the accumulated text as JSON on each
|
||||
/// text delta. Yields `ObjectStreamEvent::Partial` when a new valid partial parse is
|
||||
/// obtained, `ObjectStreamEvent::Delta` for every raw stream event, and
|
||||
/// `ObjectStreamEvent::Complete` when the stream finishes with the final parsed object.
|
||||
///
|
||||
/// # Errors
|
||||
///
|
||||
/// Returns `SdkError::Configuration` if both `prompt` and `messages` are set,
|
||||
/// `SdkError::NoObjectGenerated` if the final accumulated text is not valid JSON,
|
||||
/// or any provider error encountered during streaming.
|
||||
pub async fn stream_object(
|
||||
params: GenerateParams,
|
||||
schema: serde_json::Value,
|
||||
) -> Result<ObjectStream, SdkError> {
|
||||
let params = GenerateParams {
|
||||
response_format: Some(ResponseFormat {
|
||||
kind: ResponseFormatType::JsonSchema,
|
||||
json_schema: Some(schema),
|
||||
strict: true,
|
||||
}),
|
||||
..params
|
||||
};
|
||||
|
||||
let inner_stream = stream(params).await?;
|
||||
|
||||
let mapped = inner_stream.scan(
|
||||
(String::new(), Option::<serde_json::Value>::None),
|
||||
|(accumulated_text, last_parsed), event| {
|
||||
let mut events: Vec<Result<ObjectStreamEvent, SdkError>> = Vec::new();
|
||||
|
||||
match &event {
|
||||
Ok(stream_event) => {
|
||||
// Accumulate text from TextDelta events
|
||||
if let StreamEvent::TextDelta { delta, .. } = stream_event {
|
||||
accumulated_text.push_str(delta);
|
||||
|
||||
// Try incremental JSON parse
|
||||
if let Ok(parsed) = serde_json::from_str::<serde_json::Value>(accumulated_text) {
|
||||
if last_parsed.as_ref() != Some(&parsed) {
|
||||
*last_parsed = Some(parsed.clone());
|
||||
events.push(Ok(ObjectStreamEvent::Partial { object: parsed }));
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// On Finish, yield the Complete event with final parsed object
|
||||
if let StreamEvent::Finish { response, .. } = stream_event {
|
||||
match serde_json::from_str::<serde_json::Value>(accumulated_text) {
|
||||
Ok(final_object) => {
|
||||
events.push(Ok(ObjectStreamEvent::Complete {
|
||||
object: final_object,
|
||||
response: response.clone(),
|
||||
}));
|
||||
}
|
||||
Err(e) => {
|
||||
events.push(Err(SdkError::NoObjectGenerated {
|
||||
message: format!("Failed to parse final response as JSON: {e}"),
|
||||
}));
|
||||
}
|
||||
}
|
||||
} else {
|
||||
// Yield the raw delta event
|
||||
events.push(Ok(ObjectStreamEvent::Delta {
|
||||
event: stream_event.clone(),
|
||||
}));
|
||||
}
|
||||
}
|
||||
Err(e) => {
|
||||
events.push(Err(SdkError::Stream {
|
||||
message: format!("{e}"),
|
||||
}));
|
||||
}
|
||||
}
|
||||
|
||||
futures::future::ready(Some(futures::stream::iter(events)))
|
||||
},
|
||||
);
|
||||
|
||||
Ok(Box::pin(mapped.flatten()))
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
|
@ -774,4 +996,302 @@ mod tests {
|
|||
SdkError::NoObjectGenerated { .. }
|
||||
));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn generate_stop_when_halts_tool_loop() {
|
||||
let call_count = Arc::new(AtomicU32::new(0));
|
||||
let provider: Arc<dyn ProviderAdapter> = Arc::new(ToolCallMockProvider {
|
||||
call_count: call_count.clone(),
|
||||
});
|
||||
|
||||
let mut providers: HashMap<String, Arc<dyn ProviderAdapter>> = HashMap::new();
|
||||
providers.insert("mock".to_string(), provider);
|
||||
let client = Arc::new(Client::new(
|
||||
providers,
|
||||
Some("mock".to_string()),
|
||||
vec![],
|
||||
));
|
||||
|
||||
let result = generate(
|
||||
GenerateParams::new("mock-model")
|
||||
.prompt("What's the weather in SF?")
|
||||
.tools(vec![Tool::active(
|
||||
"get_weather",
|
||||
"Get weather",
|
||||
serde_json::json!({"type": "object", "properties": {"city": {"type": "string"}}}),
|
||||
|args| async move {
|
||||
let city = args["city"].as_str().unwrap_or("unknown");
|
||||
Ok(serde_json::json!(format!("72F in {}", city)))
|
||||
},
|
||||
)])
|
||||
.max_tool_rounds(5)
|
||||
.stop_when(|_steps| true) // Stop immediately after first round
|
||||
.client(client),
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
// stop_when returned true, so the tool loop should stop after 1 step
|
||||
assert_eq!(result.steps.len(), 1);
|
||||
assert_eq!(call_count.load(Ordering::SeqCst), 1);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn generate_params_builder_methods() {
|
||||
let params = GenerateParams::new("test-model")
|
||||
.prompt("hello")
|
||||
.system("you are helpful")
|
||||
.temperature(0.7)
|
||||
.top_p(0.9)
|
||||
.max_tokens(100)
|
||||
.stop_sequences(vec!["STOP".to_string()])
|
||||
.reasoning_effort("high")
|
||||
.provider("anthropic")
|
||||
.provider_options(serde_json::json!({"key": "value"}))
|
||||
.max_retries(5)
|
||||
.tool_choice(ToolChoice::Required)
|
||||
.response_format(ResponseFormat {
|
||||
kind: ResponseFormatType::JsonObject,
|
||||
json_schema: None,
|
||||
strict: false,
|
||||
})
|
||||
.max_tool_rounds(3);
|
||||
|
||||
assert_eq!(params.model, "test-model");
|
||||
assert_eq!(params.prompt.as_deref(), Some("hello"));
|
||||
assert_eq!(params.system.as_deref(), Some("you are helpful"));
|
||||
assert_eq!(params.temperature, Some(0.7));
|
||||
assert_eq!(params.top_p, Some(0.9));
|
||||
assert_eq!(params.max_tokens, Some(100));
|
||||
assert_eq!(
|
||||
params.stop_sequences,
|
||||
Some(vec!["STOP".to_string()])
|
||||
);
|
||||
assert_eq!(params.reasoning_effort.as_deref(), Some("high"));
|
||||
assert_eq!(params.provider.as_deref(), Some("anthropic"));
|
||||
assert!(params.provider_options.is_some());
|
||||
assert_eq!(params.max_retries, 5);
|
||||
assert_eq!(params.tool_choice, Some(ToolChoice::Required));
|
||||
assert!(params.response_format.is_some());
|
||||
assert_eq!(params.max_tool_rounds, 3);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn generate_params_timeout_builder() {
|
||||
let params = GenerateParams::new("test-model")
|
||||
.timeout(TimeoutConfig {
|
||||
total: Some(30.0),
|
||||
per_step: Some(10.0),
|
||||
});
|
||||
assert!(params.timeout.is_some());
|
||||
let t = params.timeout.unwrap();
|
||||
assert_eq!(t.total, Some(30.0));
|
||||
assert_eq!(t.per_step, Some(10.0));
|
||||
}
|
||||
|
||||
/// Mock provider that streams JSON tokens incrementally.
|
||||
struct StreamingJsonMockProvider {
|
||||
deltas: Vec<String>,
|
||||
full_text: String,
|
||||
}
|
||||
|
||||
impl StreamingJsonMockProvider {
|
||||
fn new(deltas: Vec<&str>) -> Self {
|
||||
let full_text: String = deltas.iter().copied().collect();
|
||||
Self {
|
||||
deltas: deltas.into_iter().map(String::from).collect(),
|
||||
full_text,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait::async_trait]
|
||||
impl ProviderAdapter for StreamingJsonMockProvider {
|
||||
fn name(&self) -> &str {
|
||||
"mock"
|
||||
}
|
||||
|
||||
async fn complete(&self, _request: &Request) -> Result<Response, SdkError> {
|
||||
Ok(Response {
|
||||
id: "resp_1".into(),
|
||||
model: "mock-model".into(),
|
||||
provider: "mock".into(),
|
||||
message: Message::assistant(&self.full_text),
|
||||
finish_reason: FinishReason::Stop,
|
||||
usage: Usage::default(),
|
||||
raw: None,
|
||||
warnings: vec![],
|
||||
rate_limit: None,
|
||||
})
|
||||
}
|
||||
|
||||
async fn stream(
|
||||
&self,
|
||||
_request: &Request,
|
||||
) -> Result<StreamEventStream, SdkError> {
|
||||
let mut events: Vec<Result<StreamEvent, SdkError>> = self
|
||||
.deltas
|
||||
.iter()
|
||||
.map(|d| Ok(StreamEvent::text_delta(d.as_str(), Some("t1".into()))))
|
||||
.collect();
|
||||
|
||||
events.push(Ok(StreamEvent::finish(
|
||||
FinishReason::Stop,
|
||||
Usage {
|
||||
input_tokens: 10,
|
||||
output_tokens: 20,
|
||||
total_tokens: 30,
|
||||
..Default::default()
|
||||
},
|
||||
Response {
|
||||
id: "resp_1".into(),
|
||||
model: "mock-model".into(),
|
||||
provider: "mock".into(),
|
||||
message: Message::assistant(&self.full_text),
|
||||
finish_reason: FinishReason::Stop,
|
||||
usage: Usage {
|
||||
input_tokens: 10,
|
||||
output_tokens: 20,
|
||||
total_tokens: 30,
|
||||
..Default::default()
|
||||
},
|
||||
raw: None,
|
||||
warnings: vec![],
|
||||
rate_limit: None,
|
||||
},
|
||||
)));
|
||||
|
||||
Ok(Box::pin(stream::iter(events)))
|
||||
}
|
||||
}
|
||||
|
||||
fn streaming_json_mock_client(deltas: Vec<&str>) -> Arc<Client> {
|
||||
let mut providers: HashMap<String, Arc<dyn ProviderAdapter>> = HashMap::new();
|
||||
providers.insert(
|
||||
"mock".to_string(),
|
||||
Arc::new(StreamingJsonMockProvider::new(deltas)),
|
||||
);
|
||||
Arc::new(Client::new(
|
||||
providers,
|
||||
Some("mock".to_string()),
|
||||
vec![],
|
||||
))
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn stream_object_yields_complete_event() {
|
||||
let client = streaming_json_mock_client(vec![r#"{"name": "Alice", "age": 30}"#]);
|
||||
|
||||
let schema = serde_json::json!({
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"name": {"type": "string"},
|
||||
"age": {"type": "integer"}
|
||||
},
|
||||
"required": ["name", "age"]
|
||||
});
|
||||
|
||||
let obj_stream = stream_object(
|
||||
GenerateParams::new("mock-model")
|
||||
.prompt("Extract info")
|
||||
.client(client),
|
||||
schema,
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
let events: Vec<ObjectStreamEvent> = obj_stream
|
||||
.filter_map(|r| futures::future::ready(r.ok()))
|
||||
.collect()
|
||||
.await;
|
||||
|
||||
let complete = events
|
||||
.iter()
|
||||
.find(|e| matches!(e, ObjectStreamEvent::Complete { .. }));
|
||||
assert!(complete.is_some(), "Expected a Complete event");
|
||||
|
||||
if let ObjectStreamEvent::Complete { object, .. } = complete.unwrap() {
|
||||
assert_eq!(object["name"], "Alice");
|
||||
assert_eq!(object["age"], 30);
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn stream_object_yields_partial_events_incrementally() {
|
||||
let client = streaming_json_mock_client(vec![
|
||||
r#"{"name""#,
|
||||
r#": "Bob""#,
|
||||
r#", "age": 25}"#,
|
||||
]);
|
||||
|
||||
let schema = serde_json::json!({
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"name": {"type": "string"},
|
||||
"age": {"type": "integer"}
|
||||
}
|
||||
});
|
||||
|
||||
let obj_stream = stream_object(
|
||||
GenerateParams::new("mock-model")
|
||||
.prompt("Extract info")
|
||||
.client(client),
|
||||
schema,
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
let events: Vec<ObjectStreamEvent> = obj_stream
|
||||
.filter_map(|r| futures::future::ready(r.ok()))
|
||||
.collect()
|
||||
.await;
|
||||
|
||||
let partial_count = events
|
||||
.iter()
|
||||
.filter(|e| matches!(e, ObjectStreamEvent::Partial { .. }))
|
||||
.count();
|
||||
|
||||
assert!(
|
||||
partial_count >= 1,
|
||||
"Expected at least one Partial event, got {partial_count}"
|
||||
);
|
||||
|
||||
let delta_count = events
|
||||
.iter()
|
||||
.filter(|e| matches!(e, ObjectStreamEvent::Delta { .. }))
|
||||
.count();
|
||||
|
||||
assert_eq!(delta_count, 3);
|
||||
|
||||
let last_complete = events
|
||||
.iter()
|
||||
.rev()
|
||||
.find(|e| matches!(e, ObjectStreamEvent::Complete { .. }));
|
||||
assert!(last_complete.is_some(), "Expected a Complete event");
|
||||
if let ObjectStreamEvent::Complete { object, .. } = last_complete.unwrap() {
|
||||
assert_eq!(object["name"], "Bob");
|
||||
assert_eq!(object["age"], 25);
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn stream_object_errors_on_invalid_final_json() {
|
||||
let client = streaming_json_mock_client(vec![r#"{"name": "Alice"#]);
|
||||
|
||||
let schema = serde_json::json!({"type": "object"});
|
||||
|
||||
let obj_stream = stream_object(
|
||||
GenerateParams::new("mock-model")
|
||||
.prompt("Extract info")
|
||||
.client(client),
|
||||
schema,
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
let results: Vec<Result<ObjectStreamEvent, SdkError>> = obj_stream.collect().await;
|
||||
|
||||
let has_error = results.iter().any(|r| r.is_err());
|
||||
assert!(has_error, "Expected an error for invalid final JSON");
|
||||
}
|
||||
}
|
||||
|
|
|
|||
File diff suppressed because it is too large
Load diff
|
|
@ -1,5 +1,7 @@
|
|||
use base64::{engine::general_purpose::STANDARD as BASE64_STANDARD, Engine};
|
||||
|
||||
use crate::error::{error_from_status_code, SdkError};
|
||||
use crate::types::{Message, Role};
|
||||
use crate::types::{Message, RateLimitInfo, Role};
|
||||
|
||||
#[derive(serde::Serialize)]
|
||||
pub struct ApiMessage {
|
||||
|
|
@ -8,6 +10,7 @@ pub struct ApiMessage {
|
|||
}
|
||||
|
||||
/// Parse an error response body, extracting the message and error code.
|
||||
///
|
||||
/// `error_code_field` is the JSON field name for the error code (e.g. "type" or "status").
|
||||
#[must_use]
|
||||
pub fn parse_error_body(
|
||||
|
|
@ -33,7 +36,9 @@ pub fn parse_error_body(
|
|||
)
|
||||
}
|
||||
|
||||
/// Send an HTTP request and read the response body, returning an error on non-success status.
|
||||
/// Send an HTTP request and read the response body.
|
||||
///
|
||||
/// Returns an error on non-success status.
|
||||
///
|
||||
/// # Errors
|
||||
///
|
||||
|
|
@ -43,14 +48,20 @@ pub async fn send_and_read_body(
|
|||
provider: &str,
|
||||
error_code_field: &str,
|
||||
) -> Result<String, SdkError> {
|
||||
let http_resp = request
|
||||
.send()
|
||||
.await
|
||||
.map_err(|e| SdkError::Network {
|
||||
message: e.to_string(),
|
||||
})?;
|
||||
let http_resp = request.send().await.map_err(|e| {
|
||||
if e.is_timeout() {
|
||||
SdkError::RequestTimeout {
|
||||
message: format!("{provider}: {e}"),
|
||||
}
|
||||
} else {
|
||||
SdkError::Network {
|
||||
message: e.to_string(),
|
||||
}
|
||||
}
|
||||
})?;
|
||||
|
||||
let status = http_resp.status();
|
||||
let retry_after = parse_retry_after(http_resp.headers());
|
||||
let body = http_resp
|
||||
.text()
|
||||
.await
|
||||
|
|
@ -66,21 +77,24 @@ pub async fn send_and_read_body(
|
|||
provider.to_string(),
|
||||
code,
|
||||
raw,
|
||||
None,
|
||||
retry_after,
|
||||
));
|
||||
}
|
||||
|
||||
Ok(body)
|
||||
}
|
||||
|
||||
/// Extract system messages from a message list, returning the joined system prompt
|
||||
/// and the remaining non-system messages.
|
||||
/// Extract system and developer messages from a message list.
|
||||
///
|
||||
/// Returns the joined system prompt and the remaining messages.
|
||||
/// Per spec, Developer role messages are merged with system messages
|
||||
/// for Anthropic and Gemini.
|
||||
#[must_use]
|
||||
pub fn extract_system_prompt(messages: &[Message]) -> (Option<String>, Vec<&Message>) {
|
||||
let mut system_parts = Vec::new();
|
||||
let mut other = Vec::new();
|
||||
for msg in messages {
|
||||
if msg.role == Role::System {
|
||||
if msg.role == Role::System || msg.role == Role::Developer {
|
||||
system_parts.push(msg.text());
|
||||
} else {
|
||||
other.push(msg);
|
||||
|
|
@ -93,3 +107,257 @@ pub fn extract_system_prompt(messages: &[Message]) -> (Option<String>, Vec<&Mess
|
|||
};
|
||||
(system, other)
|
||||
}
|
||||
|
||||
/// Check if a URL string looks like a local file path.
|
||||
#[must_use]
|
||||
pub fn is_file_path(url: &str) -> bool {
|
||||
url.starts_with('/') || url.starts_with("./") || url.starts_with("~/")
|
||||
}
|
||||
|
||||
/// Infer MIME type from a file extension.
|
||||
#[must_use]
|
||||
pub fn mime_from_extension(path: &str) -> &str {
|
||||
match path.rsplit('.').next().map(str::to_lowercase).as_deref() {
|
||||
Some("png") => "image/png",
|
||||
Some("jpg" | "jpeg") => "image/jpeg",
|
||||
Some("gif") => "image/gif",
|
||||
Some("webp") => "image/webp",
|
||||
Some("heic") => "image/heic",
|
||||
Some("heif") => "image/heif",
|
||||
Some("pdf") => "application/pdf",
|
||||
Some("wav") => "audio/wav",
|
||||
Some("mp3") => "audio/mp3",
|
||||
_ => "application/octet-stream",
|
||||
}
|
||||
}
|
||||
|
||||
/// Load a local file, returning (`base64_data`, `mime_type`).
|
||||
/// Expands ~ to home directory.
|
||||
///
|
||||
/// # Errors
|
||||
/// Returns an error if the file cannot be read.
|
||||
pub fn load_file_as_base64(path: &str) -> Result<(String, String), std::io::Error> {
|
||||
let expanded = path.strip_prefix("~/").map_or_else(
|
||||
|| path.to_string(),
|
||||
|rest| {
|
||||
let home = std::env::var("HOME").unwrap_or_else(|_| "/".to_string());
|
||||
format!("{home}/{rest}")
|
||||
},
|
||||
);
|
||||
let data = std::fs::read(&expanded)?;
|
||||
let mime = mime_from_extension(&expanded).to_string();
|
||||
let b64 = BASE64_STANDARD.encode(&data);
|
||||
Ok((b64, mime))
|
||||
}
|
||||
|
||||
/// Extract the `Retry-After` header value from an HTTP response as seconds.
|
||||
#[must_use]
|
||||
pub fn parse_retry_after(headers: &reqwest::header::HeaderMap) -> Option<f64> {
|
||||
headers
|
||||
.get("retry-after")
|
||||
.and_then(|v| v.to_str().ok())
|
||||
.and_then(|s| s.parse::<f64>().ok())
|
||||
}
|
||||
|
||||
/// Parse `x-ratelimit-*` headers into a `RateLimitInfo`.
|
||||
///
|
||||
/// Returns `None` if no rate limit headers are present.
|
||||
#[must_use]
|
||||
pub fn parse_rate_limit_headers(headers: &reqwest::header::HeaderMap) -> Option<RateLimitInfo> {
|
||||
fn header_i64(headers: &reqwest::header::HeaderMap, name: &str) -> Option<i64> {
|
||||
headers
|
||||
.get(name)
|
||||
.and_then(|v| v.to_str().ok())
|
||||
.and_then(|s| s.parse::<i64>().ok())
|
||||
}
|
||||
|
||||
fn header_str(headers: &reqwest::header::HeaderMap, name: &str) -> Option<String> {
|
||||
headers
|
||||
.get(name)
|
||||
.and_then(|v| v.to_str().ok())
|
||||
.map(String::from)
|
||||
}
|
||||
|
||||
let requests_remaining = header_i64(headers, "x-ratelimit-remaining-requests");
|
||||
let requests_limit = header_i64(headers, "x-ratelimit-limit-requests");
|
||||
let tokens_remaining = header_i64(headers, "x-ratelimit-remaining-tokens");
|
||||
let tokens_limit = header_i64(headers, "x-ratelimit-limit-tokens");
|
||||
let reset_at = header_str(headers, "x-ratelimit-reset-requests")
|
||||
.or_else(|| header_str(headers, "x-ratelimit-reset-tokens"));
|
||||
|
||||
if requests_remaining.is_none()
|
||||
&& requests_limit.is_none()
|
||||
&& tokens_remaining.is_none()
|
||||
&& tokens_limit.is_none()
|
||||
&& reset_at.is_none()
|
||||
{
|
||||
return None;
|
||||
}
|
||||
|
||||
Some(RateLimitInfo {
|
||||
requests_remaining,
|
||||
requests_limit,
|
||||
tokens_remaining,
|
||||
tokens_limit,
|
||||
reset_at,
|
||||
})
|
||||
}
|
||||
|
||||
/// Send an HTTP request, read the response body, and return it along with the response headers.
|
||||
///
|
||||
/// Returns an error on non-success status.
|
||||
///
|
||||
/// # Errors
|
||||
///
|
||||
/// Returns `SdkError::Network` on connection failure or `SdkError::Provider` on non-success status.
|
||||
pub async fn send_and_read_response(
|
||||
request: reqwest::RequestBuilder,
|
||||
provider: &str,
|
||||
error_code_field: &str,
|
||||
) -> Result<(String, reqwest::header::HeaderMap), SdkError> {
|
||||
let http_resp = request.send().await.map_err(|e| {
|
||||
if e.is_timeout() {
|
||||
SdkError::RequestTimeout {
|
||||
message: format!("{provider}: {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, error_code_field);
|
||||
return Err(error_from_status_code(
|
||||
status.as_u16(),
|
||||
msg,
|
||||
provider.to_string(),
|
||||
code,
|
||||
raw,
|
||||
retry_after,
|
||||
));
|
||||
}
|
||||
|
||||
Ok((body, headers))
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn is_file_path_absolute() {
|
||||
assert!(is_file_path("/tmp/image.png"));
|
||||
assert!(is_file_path("/home/user/photo.jpg"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn is_file_path_relative() {
|
||||
assert!(is_file_path("./image.png"));
|
||||
assert!(is_file_path("./subdir/photo.jpg"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn is_file_path_tilde() {
|
||||
assert!(is_file_path("~/image.png"));
|
||||
assert!(is_file_path("~/Documents/photo.jpg"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn is_file_path_url() {
|
||||
assert!(!is_file_path("https://example.com/image.png"));
|
||||
assert!(!is_file_path("http://example.com/image.png"));
|
||||
assert!(!is_file_path("data:image/png;base64,abc"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn mime_from_extension_known() {
|
||||
assert_eq!(mime_from_extension("photo.png"), "image/png");
|
||||
assert_eq!(mime_from_extension("photo.jpg"), "image/jpeg");
|
||||
assert_eq!(mime_from_extension("photo.jpeg"), "image/jpeg");
|
||||
assert_eq!(mime_from_extension("photo.gif"), "image/gif");
|
||||
assert_eq!(mime_from_extension("photo.webp"), "image/webp");
|
||||
assert_eq!(mime_from_extension("doc.pdf"), "application/pdf");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn mime_from_extension_unknown() {
|
||||
assert_eq!(mime_from_extension("file.xyz"), "application/octet-stream");
|
||||
assert_eq!(mime_from_extension("noext"), "application/octet-stream");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_rate_limit_headers_all_present() {
|
||||
let mut headers = reqwest::header::HeaderMap::new();
|
||||
headers.insert("x-ratelimit-remaining-requests", "99".parse().unwrap());
|
||||
headers.insert("x-ratelimit-limit-requests", "100".parse().unwrap());
|
||||
headers.insert("x-ratelimit-remaining-tokens", "9000".parse().unwrap());
|
||||
headers.insert("x-ratelimit-limit-tokens", "10000".parse().unwrap());
|
||||
headers.insert(
|
||||
"x-ratelimit-reset-requests",
|
||||
"2024-01-01T00:00:00Z".parse().unwrap(),
|
||||
);
|
||||
|
||||
let info = parse_rate_limit_headers(&headers).unwrap();
|
||||
assert_eq!(info.requests_remaining, Some(99));
|
||||
assert_eq!(info.requests_limit, Some(100));
|
||||
assert_eq!(info.tokens_remaining, Some(9000));
|
||||
assert_eq!(info.tokens_limit, Some(10000));
|
||||
assert_eq!(info.reset_at.as_deref(), Some("2024-01-01T00:00:00Z"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_rate_limit_headers_none_present() {
|
||||
let headers = reqwest::header::HeaderMap::new();
|
||||
assert!(parse_rate_limit_headers(&headers).is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_rate_limit_headers_partial() {
|
||||
let mut headers = reqwest::header::HeaderMap::new();
|
||||
headers.insert("x-ratelimit-remaining-requests", "50".parse().unwrap());
|
||||
|
||||
let info = parse_rate_limit_headers(&headers).unwrap();
|
||||
assert_eq!(info.requests_remaining, Some(50));
|
||||
assert_eq!(info.requests_limit, None);
|
||||
assert_eq!(info.tokens_remaining, None);
|
||||
assert_eq!(info.tokens_limit, None);
|
||||
assert_eq!(info.reset_at, None);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_rate_limit_headers_reset_tokens_fallback() {
|
||||
let mut headers = reqwest::header::HeaderMap::new();
|
||||
headers.insert("x-ratelimit-limit-tokens", "5000".parse().unwrap());
|
||||
headers.insert(
|
||||
"x-ratelimit-reset-tokens",
|
||||
"2024-06-01T12:00:00Z".parse().unwrap(),
|
||||
);
|
||||
|
||||
let info = parse_rate_limit_headers(&headers).unwrap();
|
||||
assert_eq!(info.tokens_limit, Some(5000));
|
||||
assert_eq!(info.reset_at.as_deref(), Some("2024-06-01T12:00:00Z"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_rate_limit_headers_invalid_values_ignored() {
|
||||
let mut headers = reqwest::header::HeaderMap::new();
|
||||
headers.insert("x-ratelimit-remaining-requests", "not-a-number".parse().unwrap());
|
||||
headers.insert("x-ratelimit-limit-tokens", "10000".parse().unwrap());
|
||||
|
||||
let info = parse_rate_limit_headers(&headers).unwrap();
|
||||
assert_eq!(info.requests_remaining, None);
|
||||
assert_eq!(info.tokens_limit, Some(10000));
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,25 +1,47 @@
|
|||
use crate::error::SdkError;
|
||||
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::provider::{ProviderAdapter, StreamEventStream};
|
||||
use crate::error::{ProviderErrorDetail, ProviderErrorKind};
|
||||
use crate::providers::common::{extract_system_prompt, send_and_read_body};
|
||||
use crate::types::{
|
||||
ContentPart, FinishReason, Message, Request, Response, Role, ToolCall, Usage,
|
||||
use crate::providers::common::{
|
||||
extract_system_prompt, parse_error_body, parse_retry_after, send_and_read_body,
|
||||
};
|
||||
use crate::types::{
|
||||
ContentPart, FinishReason, Message, Request, Response, ResponseFormat, ResponseFormatType,
|
||||
Role, StreamEvent, ToolCall, ToolChoice, ToolDefinition, Usage,
|
||||
};
|
||||
|
||||
const DEFAULT_BASE_URL: &str = "https://generativelanguage.googleapis.com/v1beta";
|
||||
|
||||
/// Provider adapter for the Google Gemini `generateContent` API.
|
||||
pub struct Adapter {
|
||||
api_key: String,
|
||||
base_url: String,
|
||||
client: reqwest::Client,
|
||||
request_timeout: std::time::Duration,
|
||||
}
|
||||
|
||||
impl Adapter {
|
||||
#[must_use]
|
||||
pub fn new(api_key: impl Into<String>) -> Self {
|
||||
let timeout = crate::types::AdapterTimeout::default();
|
||||
let client = reqwest::Client::builder()
|
||||
.connect_timeout(std::time::Duration::from_secs_f64(timeout.connect))
|
||||
.build()
|
||||
.unwrap_or_default();
|
||||
Self {
|
||||
api_key: api_key.into(),
|
||||
client: reqwest::Client::new(),
|
||||
base_url: DEFAULT_BASE_URL.to_string(),
|
||||
client,
|
||||
request_timeout: std::time::Duration::from_secs_f64(timeout.request),
|
||||
}
|
||||
}
|
||||
|
||||
#[must_use]
|
||||
pub fn with_base_url(mut self, base_url: impl Into<String>) -> Self {
|
||||
self.base_url = base_url.into();
|
||||
self
|
||||
}
|
||||
}
|
||||
|
||||
// --- Request types ---
|
||||
|
|
@ -32,22 +54,21 @@ struct ApiRequest {
|
|||
system_instruction: Option<SystemInstruction>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
generation_config: Option<GenerationConfig>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
tools: Option<Vec<GeminiToolGroup>>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
tool_config: Option<serde_json::Value>,
|
||||
}
|
||||
|
||||
#[derive(serde::Serialize)]
|
||||
struct Content {
|
||||
role: String,
|
||||
parts: Vec<Part>,
|
||||
parts: Vec<serde_json::Value>,
|
||||
}
|
||||
|
||||
#[derive(serde::Serialize)]
|
||||
struct SystemInstruction {
|
||||
parts: Vec<Part>,
|
||||
}
|
||||
|
||||
#[derive(serde::Serialize)]
|
||||
struct Part {
|
||||
text: String,
|
||||
parts: Vec<serde_json::Value>,
|
||||
}
|
||||
|
||||
#[derive(serde::Serialize)]
|
||||
|
|
@ -61,6 +82,24 @@ struct GenerationConfig {
|
|||
top_p: Option<f64>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
stop_sequences: Option<Vec<String>>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
response_mime_type: Option<String>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
response_schema: Option<serde_json::Value>,
|
||||
}
|
||||
|
||||
/// Gemini groups function declarations under a `tools` array.
|
||||
#[derive(serde::Serialize)]
|
||||
#[serde(rename_all = "camelCase")]
|
||||
struct GeminiToolGroup {
|
||||
function_declarations: Vec<GeminiFunctionDecl>,
|
||||
}
|
||||
|
||||
#[derive(serde::Serialize)]
|
||||
struct GeminiFunctionDecl {
|
||||
name: String,
|
||||
description: String,
|
||||
parameters: serde_json::Value,
|
||||
}
|
||||
|
||||
// --- Response types ---
|
||||
|
|
@ -91,9 +130,15 @@ struct UsageMetadata {
|
|||
prompt_token_count: Option<i64>,
|
||||
candidates_token_count: Option<i64>,
|
||||
total_token_count: Option<i64>,
|
||||
thoughts_token_count: Option<i64>,
|
||||
cached_content_token_count: Option<i64>,
|
||||
}
|
||||
|
||||
fn map_finish_reason(reason: Option<&str>) -> FinishReason {
|
||||
/// Map Gemini's finish reason, inferring `ToolCalls` from content when needed.
|
||||
fn map_finish_reason(reason: Option<&str>, has_function_calls: bool) -> FinishReason {
|
||||
if has_function_calls {
|
||||
return FinishReason::ToolCalls;
|
||||
}
|
||||
match reason {
|
||||
Some("STOP") | None => FinishReason::Stop,
|
||||
Some("MAX_TOKENS") => FinishReason::Length,
|
||||
|
|
@ -121,6 +166,520 @@ fn parse_part(part: &serde_json::Value) -> Option<ContentPart> {
|
|||
None
|
||||
}
|
||||
|
||||
/// Check if any parts contain function calls.
|
||||
fn parts_have_function_calls(parts: &[serde_json::Value]) -> bool {
|
||||
parts.iter().any(|p| p.get("functionCall").is_some())
|
||||
}
|
||||
|
||||
/// Build a mapping from tool call ID to function name by scanning assistant messages.
|
||||
///
|
||||
/// Gemini uses function names (not call IDs) in `functionResponse`. Since the adapter
|
||||
/// generates synthetic UUIDs as tool call IDs, we need this mapping to recover the
|
||||
/// original function name when sending tool results back.
|
||||
fn build_tool_call_id_to_name(messages: &[&Message]) -> std::collections::HashMap<String, String> {
|
||||
let mut map = std::collections::HashMap::new();
|
||||
for msg in messages {
|
||||
if msg.role == Role::Assistant {
|
||||
for part in &msg.content {
|
||||
if let ContentPart::ToolCall(tc) = part {
|
||||
map.insert(tc.id.clone(), tc.name.clone());
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
map
|
||||
}
|
||||
|
||||
/// Translate unified messages to Gemini content format.
|
||||
fn translate_messages(messages: &[&Message]) -> Vec<Content> {
|
||||
let id_to_name = build_tool_call_id_to_name(messages);
|
||||
let mut contents: Vec<Content> = Vec::new();
|
||||
|
||||
for msg in messages {
|
||||
let role = match msg.role {
|
||||
Role::Assistant => "model",
|
||||
Role::User | Role::Tool => "user",
|
||||
Role::System | Role::Developer => continue,
|
||||
};
|
||||
|
||||
let parts: Vec<serde_json::Value> = msg
|
||||
.content
|
||||
.iter()
|
||||
.filter_map(|part| match part {
|
||||
ContentPart::Text(text) => Some(serde_json::json!({"text": text})),
|
||||
ContentPart::ToolCall(tc) => Some(serde_json::json!({
|
||||
"functionCall": {
|
||||
"name": tc.name,
|
||||
"args": tc.arguments,
|
||||
}
|
||||
})),
|
||||
ContentPart::Image(img) => {
|
||||
img.url.as_ref().map_or_else(
|
||||
|| {
|
||||
img.data.as_ref().map(|data| {
|
||||
let mime = img.media_type.as_deref().unwrap_or("image/png");
|
||||
let b64 = BASE64_STANDARD.encode(data);
|
||||
serde_json::json!({"inlineData": {"mimeType": mime, "data": b64}})
|
||||
})
|
||||
},
|
||||
|url| {
|
||||
if crate::providers::common::is_file_path(url) {
|
||||
match crate::providers::common::load_file_as_base64(url) {
|
||||
Ok((b64, mime)) => Some(serde_json::json!({"inlineData": {"mimeType": mime, "data": b64}})),
|
||||
Err(_) => None,
|
||||
}
|
||||
} else {
|
||||
let mime = img.media_type.as_deref().unwrap_or("image/png");
|
||||
Some(serde_json::json!({"fileData": {"mimeType": mime, "fileUri": url}}))
|
||||
}
|
||||
},
|
||||
)
|
||||
}
|
||||
ContentPart::ToolResult(tr) => {
|
||||
// Gemini's functionResponse uses the function *name*, not the call ID.
|
||||
// Look up the original function name from the tool call mapping.
|
||||
let function_name = id_to_name
|
||||
.get(&tr.tool_call_id)
|
||||
.cloned()
|
||||
.unwrap_or_else(|| tr.tool_call_id.clone());
|
||||
let response = tr.content.as_str().map_or_else(
|
||||
|| {
|
||||
if tr.content.is_object() {
|
||||
tr.content.clone()
|
||||
} else {
|
||||
serde_json::json!({"result": tr.content.to_string()})
|
||||
}
|
||||
},
|
||||
|s| serde_json::json!({"result": s}),
|
||||
);
|
||||
Some(serde_json::json!({
|
||||
"functionResponse": {
|
||||
"name": function_name,
|
||||
"response": response,
|
||||
}
|
||||
}))
|
||||
}
|
||||
_ => None,
|
||||
})
|
||||
.collect();
|
||||
|
||||
if parts.is_empty() {
|
||||
continue;
|
||||
}
|
||||
|
||||
contents.push(Content {
|
||||
role: role.to_string(),
|
||||
parts,
|
||||
});
|
||||
}
|
||||
|
||||
contents
|
||||
}
|
||||
|
||||
/// Translate unified tool definitions to Gemini's format.
|
||||
fn translate_tools(tools: &[ToolDefinition]) -> Vec<GeminiToolGroup> {
|
||||
vec![GeminiToolGroup {
|
||||
function_declarations: tools
|
||||
.iter()
|
||||
.map(|t| GeminiFunctionDecl {
|
||||
name: t.name.clone(),
|
||||
description: t.description.clone(),
|
||||
parameters: t.parameters.clone(),
|
||||
})
|
||||
.collect(),
|
||||
}]
|
||||
}
|
||||
|
||||
/// Translate unified `ToolChoice` to Gemini's `toolConfig`.
|
||||
fn translate_tool_choice(choice: &ToolChoice) -> serde_json::Value {
|
||||
match choice {
|
||||
ToolChoice::Auto => serde_json::json!({
|
||||
"functionCallingConfig": {"mode": "AUTO"}
|
||||
}),
|
||||
ToolChoice::None => serde_json::json!({
|
||||
"functionCallingConfig": {"mode": "NONE"}
|
||||
}),
|
||||
ToolChoice::Required => serde_json::json!({
|
||||
"functionCallingConfig": {"mode": "ANY"}
|
||||
}),
|
||||
ToolChoice::Named { tool_name } => serde_json::json!({
|
||||
"functionCallingConfig": {
|
||||
"mode": "ANY",
|
||||
"allowedFunctionNames": [tool_name],
|
||||
}
|
||||
}),
|
||||
}
|
||||
}
|
||||
|
||||
/// Translate unified `ResponseFormat` to Gemini generation config fields.
|
||||
///
|
||||
/// Returns `(response_mime_type, response_schema)`.
|
||||
fn translate_response_format(
|
||||
format: &ResponseFormat,
|
||||
) -> (Option<String>, Option<serde_json::Value>) {
|
||||
match format.kind {
|
||||
ResponseFormatType::Text => (None, None),
|
||||
ResponseFormatType::JsonObject => (Some("application/json".to_string()), None),
|
||||
ResponseFormatType::JsonSchema => (
|
||||
Some("application/json".to_string()),
|
||||
format.json_schema.clone(),
|
||||
),
|
||||
}
|
||||
}
|
||||
|
||||
/// Build the Gemini API request body from a unified `Request`.
|
||||
fn build_api_request(request: &Request) -> ApiRequest {
|
||||
let (system_text, other_messages) = extract_system_prompt(&request.messages);
|
||||
|
||||
let system_instruction = system_text.map(|text| SystemInstruction {
|
||||
parts: vec![serde_json::json!({"text": text})],
|
||||
});
|
||||
|
||||
let contents = translate_messages(&other_messages);
|
||||
|
||||
let (response_mime_type, response_schema) = request
|
||||
.response_format
|
||||
.as_ref()
|
||||
.map_or((None, None), translate_response_format);
|
||||
|
||||
let generation_config = GenerationConfig {
|
||||
temperature: request.temperature,
|
||||
max_output_tokens: request.max_tokens,
|
||||
top_p: request.top_p,
|
||||
stop_sequences: request.stop_sequences.clone(),
|
||||
response_mime_type,
|
||||
response_schema,
|
||||
};
|
||||
|
||||
let api_tools = request.tools.as_ref().map(|t| translate_tools(t));
|
||||
let tool_config = request.tool_choice.as_ref().map(translate_tool_choice);
|
||||
|
||||
ApiRequest {
|
||||
contents,
|
||||
system_instruction,
|
||||
generation_config: Some(generation_config),
|
||||
tools: api_tools,
|
||||
tool_config,
|
||||
}
|
||||
}
|
||||
|
||||
/// Convert `UsageMetadata` from the Gemini API into a unified `Usage`.
|
||||
fn parse_usage(metadata: Option<&UsageMetadata>) -> Usage {
|
||||
metadata.map_or_else(Usage::default, |u| {
|
||||
let input = u.prompt_token_count.unwrap_or(0);
|
||||
let output = u.candidates_token_count.unwrap_or(0);
|
||||
let total = u.total_token_count.unwrap_or(input + output);
|
||||
Usage {
|
||||
input_tokens: input,
|
||||
output_tokens: output,
|
||||
total_tokens: total,
|
||||
reasoning_tokens: u.thoughts_token_count,
|
||||
cache_read_tokens: u.cached_content_token_count,
|
||||
..Usage::default()
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
/// 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`.
|
||||
async fn send_streaming_request(
|
||||
request: reqwest::RequestBuilder,
|
||||
) -> Result<reqwest::Response, SdkError> {
|
||||
let http_resp = request.send().await.map_err(|e| SdkError::Network {
|
||||
message: e.to_string(),
|
||||
})?;
|
||||
|
||||
let status = http_resp.status();
|
||||
if !status.is_success() {
|
||||
let retry_after = parse_retry_after(http_resp.headers());
|
||||
let body = http_resp.text().await.map_err(|e| SdkError::Network {
|
||||
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,
|
||||
));
|
||||
}
|
||||
|
||||
Ok(http_resp)
|
||||
}
|
||||
|
||||
/// Process a stream of SSE chunks from the Gemini `streamGenerateContent` endpoint
|
||||
/// and yield `StreamEvent` values.
|
||||
fn process_sse_stream(http_resp: reqwest::Response, model: String) -> StreamEventStream {
|
||||
Box::pin(stream::unfold(
|
||||
SseStreamState::new(http_resp, model),
|
||||
|mut state| async move {
|
||||
// If we have buffered events, yield them first.
|
||||
if let Some(event) = state.pending_events.pop_front() {
|
||||
return Some((Ok(event), state));
|
||||
}
|
||||
|
||||
// Read SSE lines until we get a data payload or the stream ends.
|
||||
loop {
|
||||
let line = match state.read_line().await {
|
||||
Ok(Some(line)) => line,
|
||||
Ok(None) => {
|
||||
// Stream ended. Emit Finish if we haven't yet.
|
||||
if !state.finished {
|
||||
state.finished = true;
|
||||
let event = state.build_finish_event();
|
||||
return Some((Ok(event), state));
|
||||
}
|
||||
return None;
|
||||
}
|
||||
Err(e) => return Some((Err(e), state)),
|
||||
};
|
||||
|
||||
// SSE format: lines starting with "data:" carry the payload.
|
||||
let data = if let Some(stripped) = line.strip_prefix("data:") {
|
||||
stripped.trim()
|
||||
} else {
|
||||
// Ignore non-data lines (empty lines, comments, event: lines).
|
||||
continue;
|
||||
};
|
||||
|
||||
// Skip empty data lines.
|
||||
if data.is_empty() {
|
||||
continue;
|
||||
}
|
||||
|
||||
// Parse the JSON chunk.
|
||||
let chunk: ApiResponse = match serde_json::from_str(data) {
|
||||
Ok(v) => v,
|
||||
Err(e) => {
|
||||
return Some((
|
||||
Err(SdkError::Stream {
|
||||
message: format!("failed to parse Gemini SSE chunk: {e}"),
|
||||
}),
|
||||
state,
|
||||
));
|
||||
}
|
||||
};
|
||||
|
||||
// Extract events from this chunk.
|
||||
state.process_chunk(&chunk);
|
||||
|
||||
// Track usage from every chunk; the final one will have the totals.
|
||||
if let Some(ref usage_meta) = chunk.usage_metadata {
|
||||
state.usage = parse_usage(Some(usage_meta));
|
||||
}
|
||||
|
||||
// Extract finish reason from the candidate if present.
|
||||
let candidate_finish = chunk
|
||||
.candidates
|
||||
.as_ref()
|
||||
.and_then(|c| c.first())
|
||||
.and_then(|c| c.finish_reason.clone());
|
||||
if let Some(reason) = candidate_finish {
|
||||
state.finish_reason_str = Some(reason);
|
||||
}
|
||||
|
||||
// Yield the first buffered event if any were produced.
|
||||
if let Some(event) = state.pending_events.pop_front() {
|
||||
return Some((Ok(event), state));
|
||||
}
|
||||
// If no events were produced from this chunk, continue reading.
|
||||
}
|
||||
},
|
||||
))
|
||||
}
|
||||
|
||||
/// Internal state for the SSE stream processor.
|
||||
struct SseStreamState {
|
||||
http_resp: reqwest::Response,
|
||||
model: String,
|
||||
/// Buffered SSE text not yet split into complete lines.
|
||||
line_buffer: String,
|
||||
/// Events extracted from a chunk but not yet yielded.
|
||||
pending_events: std::collections::VecDeque<StreamEvent>,
|
||||
/// Whether we have emitted a `StreamStart` event.
|
||||
stream_started: bool,
|
||||
/// Whether we have emitted a `TextStart` event.
|
||||
text_started: bool,
|
||||
/// Accumulated text across all chunks.
|
||||
accumulated_text: String,
|
||||
/// Accumulated tool calls across all chunks.
|
||||
accumulated_tool_calls: Vec<ToolCall>,
|
||||
/// The `text_id` used for `TextStart`/`TextDelta`/`TextEnd`.
|
||||
text_id: String,
|
||||
/// Latest usage metadata (updated per chunk; final chunk has totals).
|
||||
usage: Usage,
|
||||
/// The finish reason string from the candidate, if received.
|
||||
finish_reason_str: Option<String>,
|
||||
/// Whether we have emitted the `Finish` event.
|
||||
finished: bool,
|
||||
}
|
||||
|
||||
impl SseStreamState {
|
||||
fn new(http_resp: reqwest::Response, model: String) -> Self {
|
||||
Self {
|
||||
http_resp,
|
||||
model,
|
||||
line_buffer: String::new(),
|
||||
pending_events: std::collections::VecDeque::new(),
|
||||
stream_started: false,
|
||||
text_started: false,
|
||||
accumulated_text: String::new(),
|
||||
accumulated_tool_calls: Vec::new(),
|
||||
text_id: uuid::Uuid::new_v4().to_string(),
|
||||
usage: Usage::default(),
|
||||
finish_reason_str: None,
|
||||
finished: false,
|
||||
}
|
||||
}
|
||||
|
||||
/// Read the next complete line from the HTTP byte stream.
|
||||
///
|
||||
/// Returns `Ok(None)` when the stream is exhausted.
|
||||
async fn read_line(&mut self) -> Result<Option<String>, SdkError> {
|
||||
loop {
|
||||
// Check if we already have a complete line in the buffer.
|
||||
if let Some(newline_pos) = self.line_buffer.find('\n') {
|
||||
let line = self.line_buffer[..newline_pos]
|
||||
.trim_end_matches('\r')
|
||||
.to_string();
|
||||
self.line_buffer = self.line_buffer[newline_pos + 1..].to_string();
|
||||
return Ok(Some(line));
|
||||
}
|
||||
|
||||
// Read more bytes from the HTTP response.
|
||||
match self.http_resp.chunk().await {
|
||||
Ok(Some(bytes)) => {
|
||||
let text = String::from_utf8_lossy(&bytes);
|
||||
self.line_buffer.push_str(&text);
|
||||
}
|
||||
Ok(None) => {
|
||||
// Stream ended. Return any remaining buffered content.
|
||||
if self.line_buffer.is_empty() {
|
||||
return Ok(None);
|
||||
}
|
||||
let remaining = std::mem::take(&mut self.line_buffer);
|
||||
let line = remaining.trim_end_matches('\r').to_string();
|
||||
if line.is_empty() {
|
||||
return Ok(None);
|
||||
}
|
||||
return Ok(Some(line));
|
||||
}
|
||||
Err(e) => {
|
||||
return Err(SdkError::Stream {
|
||||
message: format!("error reading Gemini stream: {e}"),
|
||||
});
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Extract stream events from a parsed SSE chunk and buffer them.
|
||||
fn process_chunk(&mut self, chunk: &ApiResponse) {
|
||||
if !self.stream_started {
|
||||
self.stream_started = true;
|
||||
self.pending_events.push_back(StreamEvent::StreamStart);
|
||||
}
|
||||
|
||||
let parts = chunk
|
||||
.candidates
|
||||
.as_ref()
|
||||
.and_then(|c| c.first())
|
||||
.and_then(|c| c.content.as_ref())
|
||||
.and_then(|c| c.parts.as_ref());
|
||||
|
||||
let Some(parts) = parts else {
|
||||
return;
|
||||
};
|
||||
|
||||
for part in parts {
|
||||
if let Some(text) = part.get("text").and_then(serde_json::Value::as_str) {
|
||||
if !self.text_started {
|
||||
self.text_started = true;
|
||||
self.pending_events.push_back(StreamEvent::TextStart {
|
||||
text_id: Some(self.text_id.clone()),
|
||||
});
|
||||
}
|
||||
self.accumulated_text.push_str(text);
|
||||
self.pending_events
|
||||
.push_back(StreamEvent::text_delta(text, Some(self.text_id.clone())));
|
||||
} else if let Some(fc) = part.get("functionCall") {
|
||||
let name = fc
|
||||
.get("name")
|
||||
.and_then(serde_json::Value::as_str)
|
||||
.unwrap_or("")
|
||||
.to_string();
|
||||
let args = fc
|
||||
.get("args")
|
||||
.cloned()
|
||||
.unwrap_or_else(|| serde_json::Value::Object(serde_json::Map::new()));
|
||||
let tool_call = ToolCall::new(uuid::Uuid::new_v4().to_string(), name, args);
|
||||
|
||||
// Gemini delivers function calls as complete objects in a single chunk.
|
||||
self.pending_events
|
||||
.push_back(StreamEvent::ToolCallStart {
|
||||
tool_call: tool_call.clone(),
|
||||
});
|
||||
self.pending_events.push_back(StreamEvent::ToolCallEnd {
|
||||
tool_call: tool_call.clone(),
|
||||
});
|
||||
self.accumulated_tool_calls.push(tool_call);
|
||||
}
|
||||
}
|
||||
|
||||
// If a finish reason is present on this chunk's candidate, emit TextEnd.
|
||||
let has_finish_reason = chunk
|
||||
.candidates
|
||||
.as_ref()
|
||||
.and_then(|c| c.first())
|
||||
.and_then(|c| c.finish_reason.as_ref())
|
||||
.is_some();
|
||||
|
||||
if has_finish_reason && self.text_started {
|
||||
self.pending_events.push_back(StreamEvent::TextEnd {
|
||||
text_id: Some(self.text_id.clone()),
|
||||
});
|
||||
}
|
||||
}
|
||||
|
||||
/// Build the final `Finish` event from accumulated state.
|
||||
fn build_finish_event(&self) -> StreamEvent {
|
||||
let has_tool_calls = !self.accumulated_tool_calls.is_empty();
|
||||
let finish_reason =
|
||||
map_finish_reason(self.finish_reason_str.as_deref(), has_tool_calls);
|
||||
|
||||
let mut content_parts: Vec<ContentPart> = Vec::new();
|
||||
if !self.accumulated_text.is_empty() {
|
||||
content_parts.push(ContentPart::text(&self.accumulated_text));
|
||||
}
|
||||
for tc in &self.accumulated_tool_calls {
|
||||
content_parts.push(ContentPart::ToolCall(tc.clone()));
|
||||
}
|
||||
|
||||
let response = Response {
|
||||
id: uuid::Uuid::new_v4().to_string(),
|
||||
model: self.model.clone(),
|
||||
provider: "gemini".to_string(),
|
||||
message: Message {
|
||||
role: Role::Assistant,
|
||||
content: content_parts,
|
||||
name: None,
|
||||
tool_call_id: None,
|
||||
},
|
||||
finish_reason: finish_reason.clone(),
|
||||
usage: self.usage.clone(),
|
||||
raw: None,
|
||||
warnings: vec![],
|
||||
rate_limit: None,
|
||||
};
|
||||
|
||||
StreamEvent::finish(finish_reason, self.usage.clone(), response)
|
||||
}
|
||||
}
|
||||
|
||||
#[allow(clippy::unnecessary_literal_bound)]
|
||||
#[async_trait::async_trait]
|
||||
impl ProviderAdapter for Adapter {
|
||||
|
|
@ -129,48 +688,15 @@ impl ProviderAdapter for Adapter {
|
|||
}
|
||||
|
||||
async fn complete(&self, request: &Request) -> Result<Response, SdkError> {
|
||||
let (system_text, other_messages) = extract_system_prompt(&request.messages);
|
||||
|
||||
let system_instruction = system_text.map(|text| SystemInstruction {
|
||||
parts: vec![Part { text }],
|
||||
});
|
||||
|
||||
let contents: Vec<Content> = other_messages
|
||||
.iter()
|
||||
.map(|msg| {
|
||||
let role = match msg.role {
|
||||
Role::Assistant => "model",
|
||||
Role::System | Role::User | Role::Tool | Role::Developer => "user",
|
||||
};
|
||||
Content {
|
||||
role: role.to_string(),
|
||||
parts: vec![Part {
|
||||
text: msg.text(),
|
||||
}],
|
||||
}
|
||||
})
|
||||
.collect();
|
||||
|
||||
let generation_config = GenerationConfig {
|
||||
temperature: request.temperature,
|
||||
max_output_tokens: request.max_tokens,
|
||||
top_p: request.top_p,
|
||||
stop_sequences: request.stop_sequences.clone(),
|
||||
};
|
||||
|
||||
let api_request = ApiRequest {
|
||||
contents,
|
||||
system_instruction,
|
||||
generation_config: Some(generation_config),
|
||||
};
|
||||
let api_request = build_api_request(request);
|
||||
|
||||
let url = format!(
|
||||
"https://generativelanguage.googleapis.com/v1beta/models/{}:generateContent?key={}",
|
||||
request.model, self.api_key
|
||||
"{}/models/{}:generateContent?key={}",
|
||||
self.base_url, request.model, self.api_key
|
||||
);
|
||||
|
||||
let body = send_and_read_body(
|
||||
self.client.post(&url).json(&api_request),
|
||||
self.client.post(&url).json(&api_request).timeout(self.request_timeout),
|
||||
"gemini",
|
||||
"status",
|
||||
)
|
||||
|
|
@ -187,32 +713,24 @@ impl ProviderAdapter for Adapter {
|
|||
.and_then(|c| c.first())
|
||||
.ok_or_else(|| SdkError::Provider {
|
||||
kind: ProviderErrorKind::Server,
|
||||
detail: Box::new(ProviderErrorDetail::new("no candidates in Gemini response", "gemini")),
|
||||
detail: Box::new(ProviderErrorDetail::new(
|
||||
"no candidates in Gemini response",
|
||||
"gemini",
|
||||
)),
|
||||
})?;
|
||||
|
||||
let content_parts: Vec<ContentPart> = candidate
|
||||
.content
|
||||
.as_ref()
|
||||
.and_then(|c| c.parts.as_ref())
|
||||
let raw_parts = candidate.content.as_ref().and_then(|c| c.parts.as_ref());
|
||||
|
||||
let content_parts: Vec<ContentPart> = raw_parts
|
||||
.map(|parts| parts.iter().filter_map(parse_part).collect())
|
||||
.unwrap_or_default();
|
||||
|
||||
let finish_reason = map_finish_reason(candidate.finish_reason.as_deref());
|
||||
// Gemini has no dedicated tool_calls finish reason; infer from parts
|
||||
let has_tool_calls = raw_parts.is_some_and(|p| parts_have_function_calls(p));
|
||||
let finish_reason =
|
||||
map_finish_reason(candidate.finish_reason.as_deref(), has_tool_calls);
|
||||
|
||||
let usage = api_resp
|
||||
.usage_metadata
|
||||
.as_ref()
|
||||
.map_or_else(Usage::default, |u| {
|
||||
let input = u.prompt_token_count.unwrap_or(0);
|
||||
let output = u.candidates_token_count.unwrap_or(0);
|
||||
let total = u.total_token_count.unwrap_or(input + output);
|
||||
Usage {
|
||||
input_tokens: input,
|
||||
output_tokens: output,
|
||||
total_tokens: total,
|
||||
..Usage::default()
|
||||
}
|
||||
});
|
||||
let usage = parse_usage(api_resp.usage_metadata.as_ref());
|
||||
|
||||
Ok(Response {
|
||||
id: uuid::Uuid::new_v4().to_string(),
|
||||
|
|
@ -232,9 +750,17 @@ impl ProviderAdapter for Adapter {
|
|||
})
|
||||
}
|
||||
|
||||
async fn stream(&self, _request: &Request) -> Result<StreamEventStream, SdkError> {
|
||||
Err(SdkError::Configuration {
|
||||
message: "streaming not yet implemented".to_string(),
|
||||
})
|
||||
async fn stream(&self, request: &Request) -> Result<StreamEventStream, SdkError> {
|
||||
let api_request = build_api_request(request);
|
||||
|
||||
let url = format!(
|
||||
"{}/models/{}:streamGenerateContent?alt=sse&key={}",
|
||||
self.base_url, request.model, self.api_key
|
||||
);
|
||||
|
||||
let http_resp =
|
||||
send_streaming_request(self.client.post(&url).json(&api_request)).await?;
|
||||
|
||||
Ok(process_sse_stream(http_resp, request.model.clone()))
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -2,7 +2,9 @@ pub mod anthropic;
|
|||
pub mod common;
|
||||
pub mod gemini;
|
||||
pub mod openai;
|
||||
pub mod openai_compatible;
|
||||
|
||||
pub use anthropic::Adapter as AnthropicAdapter;
|
||||
pub use gemini::Adapter as GeminiAdapter;
|
||||
pub use openai::Adapter as OpenAiAdapter;
|
||||
pub use openai_compatible::Adapter as OpenAiCompatibleAdapter;
|
||||
|
|
|
|||
|
|
@ -1,95 +1,741 @@
|
|||
use base64::{engine::general_purpose::STANDARD as BASE64_STANDARD, Engine};
|
||||
use futures::StreamExt;
|
||||
|
||||
use crate::error::SdkError;
|
||||
use crate::provider::{ProviderAdapter, StreamEventStream};
|
||||
use crate::error::{ProviderErrorDetail, ProviderErrorKind};
|
||||
use crate::providers::common::{send_and_read_body, ApiMessage};
|
||||
use crate::providers::common::{
|
||||
parse_error_body, parse_rate_limit_headers, parse_retry_after, send_and_read_response,
|
||||
};
|
||||
use crate::types::{
|
||||
ContentPart, FinishReason, Message, Request, Response, Role, ToolCall, Usage,
|
||||
ContentPart, FinishReason, Message, Request, Response, ResponseFormat, ResponseFormatType,
|
||||
Role, StreamEvent, ToolCall, ToolChoice, ToolDefinition, Usage,
|
||||
};
|
||||
|
||||
/// Provider adapter for the `OpenAI` Chat Completions API.
|
||||
/// Provider adapter for the `OpenAI` Responses API (`/v1/responses`).
|
||||
///
|
||||
/// Per spec Section 2.7, this adapter uses the Responses API (not Chat Completions)
|
||||
/// to properly surface reasoning tokens, built-in tools, and server-side state.
|
||||
pub struct Adapter {
|
||||
api_key: String,
|
||||
base_url: String,
|
||||
client: reqwest::Client,
|
||||
request_timeout: std::time::Duration,
|
||||
}
|
||||
|
||||
impl Adapter {
|
||||
#[must_use]
|
||||
pub fn new(api_key: impl Into<String>) -> Self {
|
||||
let timeout = crate::types::AdapterTimeout::default();
|
||||
let client = reqwest::Client::builder()
|
||||
.connect_timeout(std::time::Duration::from_secs_f64(timeout.connect))
|
||||
.build()
|
||||
.unwrap_or_default();
|
||||
Self {
|
||||
api_key: api_key.into(),
|
||||
client: reqwest::Client::new(),
|
||||
base_url: "https://api.openai.com/v1".to_string(),
|
||||
client,
|
||||
request_timeout: std::time::Duration::from_secs_f64(timeout.request),
|
||||
}
|
||||
}
|
||||
|
||||
#[must_use]
|
||||
pub fn with_base_url(mut self, base_url: impl Into<String>) -> Self {
|
||||
self.base_url = base_url.into();
|
||||
self
|
||||
}
|
||||
}
|
||||
|
||||
// --- Request types ---
|
||||
// --- Request types (Responses API format) ---
|
||||
|
||||
#[derive(serde::Serialize)]
|
||||
struct ApiRequest {
|
||||
model: String,
|
||||
messages: Vec<ApiMessage>,
|
||||
input: Vec<serde_json::Value>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
instructions: Option<String>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
temperature: Option<f64>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
max_tokens: Option<i64>,
|
||||
max_output_tokens: Option<i64>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
top_p: Option<f64>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
stop: Option<Vec<String>>,
|
||||
tools: Option<Vec<serde_json::Value>>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
tool_choice: Option<serde_json::Value>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
reasoning: Option<serde_json::Value>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
text: Option<serde_json::Value>,
|
||||
#[serde(skip_serializing_if = "std::ops::Not::not")]
|
||||
stream: bool,
|
||||
}
|
||||
|
||||
// --- Response types ---
|
||||
// --- Response types (Responses API format) ---
|
||||
|
||||
#[derive(serde::Deserialize)]
|
||||
struct ApiResponse {
|
||||
id: String,
|
||||
model: String,
|
||||
choices: Vec<ApiChoice>,
|
||||
model: Option<String>,
|
||||
output: Vec<serde_json::Value>,
|
||||
status: Option<String>,
|
||||
usage: Option<ApiUsage>,
|
||||
}
|
||||
|
||||
#[derive(serde::Deserialize)]
|
||||
struct ApiChoice {
|
||||
message: ApiChoiceMessage,
|
||||
finish_reason: Option<String>,
|
||||
}
|
||||
|
||||
#[derive(serde::Deserialize)]
|
||||
struct ApiChoiceMessage {
|
||||
content: Option<String>,
|
||||
tool_calls: Option<Vec<ApiToolCall>>,
|
||||
}
|
||||
|
||||
#[derive(serde::Deserialize)]
|
||||
struct ApiToolCall {
|
||||
id: String,
|
||||
function: ApiFunction,
|
||||
}
|
||||
|
||||
#[derive(serde::Deserialize)]
|
||||
struct ApiFunction {
|
||||
name: String,
|
||||
arguments: String,
|
||||
}
|
||||
|
||||
#[derive(serde::Deserialize)]
|
||||
#[allow(clippy::struct_field_names)]
|
||||
struct ApiUsage {
|
||||
prompt_tokens: i64,
|
||||
completion_tokens: i64,
|
||||
total_tokens: i64,
|
||||
input_tokens: i64,
|
||||
output_tokens: i64,
|
||||
total_tokens: Option<i64>,
|
||||
output_tokens_details: Option<OutputTokenDetails>,
|
||||
input_tokens_details: Option<InputTokenDetails>,
|
||||
}
|
||||
|
||||
fn map_finish_reason(reason: Option<&str>) -> FinishReason {
|
||||
match reason {
|
||||
Some("stop") | None => FinishReason::Stop,
|
||||
Some("length") => FinishReason::Length,
|
||||
Some("tool_calls") => FinishReason::ToolCalls,
|
||||
Some("content_filter") => FinishReason::ContentFilter,
|
||||
#[derive(serde::Deserialize)]
|
||||
struct OutputTokenDetails {
|
||||
reasoning_tokens: Option<i64>,
|
||||
}
|
||||
|
||||
#[derive(serde::Deserialize)]
|
||||
struct InputTokenDetails {
|
||||
cached_tokens: Option<i64>,
|
||||
}
|
||||
|
||||
/// Map the Responses API status to a `FinishReason`.
|
||||
fn map_finish_reason(status: Option<&str>, has_tool_calls: bool) -> FinishReason {
|
||||
if has_tool_calls {
|
||||
return FinishReason::ToolCalls;
|
||||
}
|
||||
match status {
|
||||
Some("completed") | None => FinishReason::Stop,
|
||||
Some("incomplete") => FinishReason::Length,
|
||||
Some("failed") => FinishReason::Error,
|
||||
Some(other) => FinishReason::Other(other.to_string()),
|
||||
}
|
||||
}
|
||||
|
||||
/// Translate unified messages to Responses API `input` array format.
|
||||
fn translate_input(messages: &[Message]) -> (Option<String>, Vec<serde_json::Value>) {
|
||||
let mut instructions_parts: Vec<String> = Vec::new();
|
||||
let mut input: Vec<serde_json::Value> = Vec::new();
|
||||
|
||||
for msg in messages {
|
||||
match msg.role {
|
||||
Role::System | Role::Developer => {
|
||||
instructions_parts.push(msg.text());
|
||||
}
|
||||
Role::User => {
|
||||
let content: Vec<serde_json::Value> = msg
|
||||
.content
|
||||
.iter()
|
||||
.filter_map(|part| match part {
|
||||
ContentPart::Text(text) => {
|
||||
Some(serde_json::json!({"type": "input_text", "text": text}))
|
||||
}
|
||||
ContentPart::Image(img) => {
|
||||
img.url.as_ref().map_or_else(
|
||||
|| {
|
||||
img.data.as_ref().map(|data| {
|
||||
let mime = img.media_type.as_deref().unwrap_or("image/png");
|
||||
let b64 = BASE64_STANDARD.encode(data);
|
||||
serde_json::json!({"type": "input_image", "image_url": format!("data:{mime};base64,{b64}")})
|
||||
})
|
||||
},
|
||||
|url| {
|
||||
if crate::providers::common::is_file_path(url) {
|
||||
match crate::providers::common::load_file_as_base64(url) {
|
||||
Ok((b64, mime)) => Some(serde_json::json!({"type": "input_image", "image_url": format!("data:{mime};base64,{b64}")})),
|
||||
Err(_) => None,
|
||||
}
|
||||
} else {
|
||||
Some(serde_json::json!({"type": "input_image", "image_url": url}))
|
||||
}
|
||||
},
|
||||
)
|
||||
}
|
||||
_ => None,
|
||||
})
|
||||
.collect();
|
||||
if !content.is_empty() {
|
||||
input.push(serde_json::json!({
|
||||
"type": "message",
|
||||
"role": "user",
|
||||
"content": content,
|
||||
}));
|
||||
}
|
||||
}
|
||||
Role::Assistant => {
|
||||
for part in &msg.content {
|
||||
match part {
|
||||
ContentPart::Text(text) => {
|
||||
input.push(serde_json::json!({
|
||||
"type": "message",
|
||||
"role": "assistant",
|
||||
"content": [{"type": "output_text", "text": text}],
|
||||
}));
|
||||
}
|
||||
ContentPart::ToolCall(tc) => {
|
||||
let args = tc
|
||||
.raw_arguments
|
||||
.as_ref()
|
||||
.map_or_else(|| tc.arguments.to_string(), Clone::clone);
|
||||
input.push(serde_json::json!({
|
||||
"type": "function_call",
|
||||
"id": tc.id,
|
||||
"call_id": tc.id,
|
||||
"name": tc.name,
|
||||
"arguments": args,
|
||||
}));
|
||||
}
|
||||
_ => {}
|
||||
}
|
||||
}
|
||||
}
|
||||
Role::Tool => {
|
||||
for part in &msg.content {
|
||||
if let ContentPart::ToolResult(tr) = part {
|
||||
let output = tr
|
||||
.content
|
||||
.as_str()
|
||||
.map_or_else(|| tr.content.to_string(), str::to_string);
|
||||
input.push(serde_json::json!({
|
||||
"type": "function_call_output",
|
||||
"call_id": tr.tool_call_id,
|
||||
"output": output,
|
||||
}));
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
let instructions = if instructions_parts.is_empty() {
|
||||
None
|
||||
} else {
|
||||
Some(instructions_parts.join("\n"))
|
||||
};
|
||||
|
||||
(instructions, input)
|
||||
}
|
||||
|
||||
/// Translate unified tool definitions to Responses API tool format.
|
||||
fn translate_tools(tools: &[ToolDefinition]) -> Vec<serde_json::Value> {
|
||||
tools
|
||||
.iter()
|
||||
.map(|t| {
|
||||
serde_json::json!({
|
||||
"type": "function",
|
||||
"name": t.name,
|
||||
"description": t.description,
|
||||
"parameters": t.parameters,
|
||||
})
|
||||
})
|
||||
.collect()
|
||||
}
|
||||
|
||||
/// Translate unified `ToolChoice` to Responses API format.
|
||||
fn translate_tool_choice(choice: &ToolChoice) -> serde_json::Value {
|
||||
match choice {
|
||||
ToolChoice::Auto => serde_json::json!("auto"),
|
||||
ToolChoice::None => serde_json::json!("none"),
|
||||
ToolChoice::Required => serde_json::json!("required"),
|
||||
ToolChoice::Named { tool_name } => {
|
||||
serde_json::json!({"type": "function", "name": tool_name})
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Translate unified `ResponseFormat` to Responses API `text` field.
|
||||
///
|
||||
/// The Responses API uses `"text": {"format": {...}}` for structured output.
|
||||
fn translate_response_format(format: &ResponseFormat) -> Option<serde_json::Value> {
|
||||
match format.kind {
|
||||
ResponseFormatType::Text => None,
|
||||
ResponseFormatType::JsonObject => {
|
||||
Some(serde_json::json!({"format": {"type": "json_object"}}))
|
||||
}
|
||||
ResponseFormatType::JsonSchema => {
|
||||
let mut schema_obj = serde_json::json!({
|
||||
"type": "json_schema",
|
||||
"name": "response",
|
||||
"strict": format.strict,
|
||||
});
|
||||
if let Some(schema) = &format.json_schema {
|
||||
schema_obj["schema"] = schema.clone();
|
||||
}
|
||||
Some(serde_json::json!({"format": schema_obj}))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Build an `ApiRequest` from a unified `Request`.
|
||||
fn build_api_request(request: &Request, stream: bool) -> ApiRequest {
|
||||
let (instructions, input) = translate_input(&request.messages);
|
||||
let api_tools = request.tools.as_ref().map(|t| translate_tools(t));
|
||||
let tool_choice = request.tool_choice.as_ref().map(translate_tool_choice);
|
||||
let reasoning = request
|
||||
.reasoning_effort
|
||||
.as_ref()
|
||||
.map(|effort| serde_json::json!({"effort": effort}));
|
||||
let text = request
|
||||
.response_format
|
||||
.as_ref()
|
||||
.and_then(translate_response_format);
|
||||
|
||||
ApiRequest {
|
||||
model: request.model.clone(),
|
||||
input,
|
||||
instructions,
|
||||
temperature: request.temperature,
|
||||
max_output_tokens: request.max_tokens,
|
||||
top_p: request.top_p,
|
||||
tools: api_tools,
|
||||
tool_choice,
|
||||
reasoning,
|
||||
text,
|
||||
stream,
|
||||
}
|
||||
}
|
||||
|
||||
/// Parse output items from the Responses API into content parts.
|
||||
fn parse_output(output: &[serde_json::Value]) -> (Vec<ContentPart>, bool) {
|
||||
let mut parts = Vec::new();
|
||||
let mut has_tool_calls = false;
|
||||
|
||||
for item in output {
|
||||
let item_type = item.get("type").and_then(serde_json::Value::as_str);
|
||||
match item_type {
|
||||
Some("message") => {
|
||||
if let Some(content) = item.get("content").and_then(|c| c.as_array()) {
|
||||
for block in content {
|
||||
if block.get("type").and_then(serde_json::Value::as_str)
|
||||
== Some("output_text")
|
||||
{
|
||||
if let Some(text) =
|
||||
block.get("text").and_then(serde_json::Value::as_str)
|
||||
{
|
||||
parts.push(ContentPart::text(text));
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
Some("function_call") => {
|
||||
has_tool_calls = true;
|
||||
let id = item
|
||||
.get("call_id")
|
||||
.or_else(|| item.get("id"))
|
||||
.and_then(serde_json::Value::as_str)
|
||||
.unwrap_or("")
|
||||
.to_string();
|
||||
let name = item
|
||||
.get("name")
|
||||
.and_then(serde_json::Value::as_str)
|
||||
.unwrap_or("")
|
||||
.to_string();
|
||||
let args_str = item
|
||||
.get("arguments")
|
||||
.and_then(serde_json::Value::as_str)
|
||||
.unwrap_or("{}");
|
||||
let arguments = serde_json::from_str(args_str)
|
||||
.unwrap_or_else(|_| serde_json::json!({}));
|
||||
let mut tc = ToolCall::new(id, name, arguments);
|
||||
tc.raw_arguments = Some(args_str.to_string());
|
||||
parts.push(ContentPart::ToolCall(tc));
|
||||
}
|
||||
_ => {}
|
||||
}
|
||||
}
|
||||
|
||||
(parts, has_tool_calls)
|
||||
}
|
||||
|
||||
// --- SSE streaming support ---
|
||||
|
||||
/// Mutable state carried through SSE stream processing.
|
||||
struct SseStreamState {
|
||||
byte_stream: std::pin::Pin<
|
||||
Box<dyn futures::Stream<Item = Result<bytes::Bytes, reqwest::Error>> + Send>,
|
||||
>,
|
||||
buffer: String,
|
||||
model: String,
|
||||
response_id: String,
|
||||
response_model: String,
|
||||
accumulated_text: String,
|
||||
tool_calls: Vec<ToolCall>,
|
||||
usage: Usage,
|
||||
finish_reason: FinishReason,
|
||||
emitted_start: bool,
|
||||
emitted_text_start: bool,
|
||||
raw_response: Option<serde_json::Value>,
|
||||
rate_limit: Option<crate::types::RateLimitInfo>,
|
||||
}
|
||||
|
||||
/// Extract complete SSE messages from the buffer.
|
||||
///
|
||||
/// Each SSE message consists of one or more lines (`event:` and `data:` prefixed)
|
||||
/// terminated by a blank line. Returns parsed (`event_type`, data) pairs.
|
||||
fn extract_sse_messages(buffer: &mut String) -> Vec<(Option<String>, String)> {
|
||||
let mut messages = Vec::new();
|
||||
|
||||
while let Some(pos) = buffer.find("\n\n") {
|
||||
let message_block = buffer[..pos].to_string();
|
||||
*buffer = buffer[pos + 2..].to_string();
|
||||
|
||||
let mut current_event: Option<String> = None;
|
||||
let mut current_data = String::new();
|
||||
|
||||
for line in message_block.lines() {
|
||||
if let Some(stripped) = line.strip_prefix("event: ") {
|
||||
current_event = Some(stripped.to_string());
|
||||
} else if let Some(stripped) = line.strip_prefix("event:") {
|
||||
current_event = Some(stripped.trim().to_string());
|
||||
} else if let Some(stripped) = line.strip_prefix("data: ") {
|
||||
if !current_data.is_empty() {
|
||||
current_data.push('\n');
|
||||
}
|
||||
current_data.push_str(stripped);
|
||||
} else if let Some(stripped) = line.strip_prefix("data:") {
|
||||
if !current_data.is_empty() {
|
||||
current_data.push('\n');
|
||||
}
|
||||
current_data.push_str(stripped.trim());
|
||||
}
|
||||
}
|
||||
|
||||
if !current_data.is_empty() {
|
||||
messages.push((current_event, current_data));
|
||||
}
|
||||
}
|
||||
|
||||
messages
|
||||
}
|
||||
|
||||
/// Dispatch SSE messages from the buffer and return the resulting `StreamEvent`s.
|
||||
fn dispatch_sse_messages(
|
||||
state: &mut SseStreamState,
|
||||
messages: Vec<(Option<String>, String)>,
|
||||
) -> Vec<StreamEvent> {
|
||||
let mut events = Vec::new();
|
||||
for (event_type, data) in messages {
|
||||
events.extend(process_sse_event(state, event_type.as_deref(), &data));
|
||||
}
|
||||
events
|
||||
}
|
||||
|
||||
/// Process the next chunk(s) from the byte stream and return `StreamEvent`s.
|
||||
async fn process_next_sse_events(
|
||||
state: &mut SseStreamState,
|
||||
) -> Result<Vec<StreamEvent>, SdkError> {
|
||||
loop {
|
||||
let messages = extract_sse_messages(&mut state.buffer);
|
||||
if !messages.is_empty() {
|
||||
return Ok(dispatch_sse_messages(state, messages));
|
||||
}
|
||||
|
||||
match state.byte_stream.next().await {
|
||||
Some(Ok(bytes)) => {
|
||||
let text = String::from_utf8_lossy(&bytes);
|
||||
state.buffer.push_str(&text);
|
||||
}
|
||||
Some(Err(e)) => {
|
||||
return Err(SdkError::Stream {
|
||||
message: e.to_string(),
|
||||
});
|
||||
}
|
||||
None => {
|
||||
// Stream ended. Process any remaining data in the buffer.
|
||||
if !state.buffer.is_empty() {
|
||||
state.buffer.push_str("\n\n");
|
||||
let messages = extract_sse_messages(&mut state.buffer);
|
||||
return Ok(dispatch_sse_messages(state, messages));
|
||||
}
|
||||
return Ok(vec![]);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Process a single SSE event and return the corresponding `StreamEvent`(s).
|
||||
fn process_sse_event(
|
||||
state: &mut SseStreamState,
|
||||
event_type: Option<&str>,
|
||||
data: &str,
|
||||
) -> Vec<StreamEvent> {
|
||||
let mut events = Vec::new();
|
||||
|
||||
if !state.emitted_start {
|
||||
state.emitted_start = true;
|
||||
events.push(StreamEvent::StreamStart);
|
||||
}
|
||||
|
||||
let json: serde_json::Value = match serde_json::from_str(data) {
|
||||
Ok(v) => v,
|
||||
Err(_) => return events,
|
||||
};
|
||||
|
||||
// Resolve event type from the `event:` SSE line or from the JSON `type` field.
|
||||
let resolved_type = event_type
|
||||
.map(str::to_string)
|
||||
.or_else(|| {
|
||||
json.get("type")
|
||||
.and_then(serde_json::Value::as_str)
|
||||
.map(str::to_string)
|
||||
})
|
||||
.unwrap_or_default();
|
||||
|
||||
match resolved_type.as_str() {
|
||||
"response.created" => handle_response_created(state, &json),
|
||||
"response.output_text.delta" => handle_text_delta(state, &json, &mut events),
|
||||
"response.function_call_arguments.delta" => {
|
||||
handle_tool_call_delta(state, &json, &mut events);
|
||||
}
|
||||
"response.output_item.done" => handle_output_item_done(state, &json, &mut events),
|
||||
"response.completed" => handle_response_completed(state, &json, &mut events),
|
||||
_ => {}
|
||||
}
|
||||
|
||||
events
|
||||
}
|
||||
|
||||
/// Handle `response.created` by extracting the response ID and model.
|
||||
fn handle_response_created(state: &mut SseStreamState, json: &serde_json::Value) {
|
||||
if let Some(id) = json
|
||||
.get("response")
|
||||
.and_then(|r| r.get("id"))
|
||||
.and_then(serde_json::Value::as_str)
|
||||
{
|
||||
state.response_id = id.to_string();
|
||||
}
|
||||
if let Some(model) = json
|
||||
.get("response")
|
||||
.and_then(|r| r.get("model"))
|
||||
.and_then(serde_json::Value::as_str)
|
||||
{
|
||||
state.response_model = model.to_string();
|
||||
}
|
||||
}
|
||||
|
||||
/// Handle `response.output_text.delta` by accumulating text and emitting events.
|
||||
fn handle_text_delta(
|
||||
state: &mut SseStreamState,
|
||||
json: &serde_json::Value,
|
||||
events: &mut Vec<StreamEvent>,
|
||||
) {
|
||||
if let Some(delta) = json.get("delta").and_then(serde_json::Value::as_str) {
|
||||
if !state.emitted_text_start {
|
||||
state.emitted_text_start = true;
|
||||
events.push(StreamEvent::TextStart { text_id: None });
|
||||
}
|
||||
state.accumulated_text.push_str(delta);
|
||||
events.push(StreamEvent::text_delta(delta, None));
|
||||
}
|
||||
}
|
||||
|
||||
/// Handle `response.function_call_arguments.delta` by accumulating args and emitting events.
|
||||
fn handle_tool_call_delta(
|
||||
state: &mut SseStreamState,
|
||||
json: &serde_json::Value,
|
||||
events: &mut Vec<StreamEvent>,
|
||||
) {
|
||||
let Some(delta) = json.get("delta").and_then(serde_json::Value::as_str) else {
|
||||
return;
|
||||
};
|
||||
|
||||
let call_id = json
|
||||
.get("call_id")
|
||||
.and_then(serde_json::Value::as_str)
|
||||
.unwrap_or("")
|
||||
.to_string();
|
||||
let item_id = json
|
||||
.get("item_id")
|
||||
.and_then(serde_json::Value::as_str)
|
||||
.unwrap_or("")
|
||||
.to_string();
|
||||
let lookup_id = if call_id.is_empty() {
|
||||
&item_id
|
||||
} else {
|
||||
&call_id
|
||||
};
|
||||
|
||||
let tc_index = state
|
||||
.tool_calls
|
||||
.iter()
|
||||
.position(|tc| tc.id == *lookup_id);
|
||||
|
||||
if let Some(idx) = tc_index {
|
||||
if let Some(ref mut raw) = state.tool_calls[idx].raw_arguments {
|
||||
raw.push_str(delta);
|
||||
}
|
||||
} else {
|
||||
let name = json
|
||||
.get("name")
|
||||
.and_then(serde_json::Value::as_str)
|
||||
.unwrap_or("")
|
||||
.to_string();
|
||||
let mut tc = ToolCall::new(lookup_id, name, serde_json::json!({}));
|
||||
tc.raw_arguments = Some(delta.to_string());
|
||||
state.tool_calls.push(tc.clone());
|
||||
events.push(StreamEvent::ToolCallStart { tool_call: tc });
|
||||
}
|
||||
|
||||
let current_tc = state
|
||||
.tool_calls
|
||||
.iter()
|
||||
.find(|tc| tc.id == *lookup_id)
|
||||
.cloned()
|
||||
.unwrap_or_else(|| ToolCall::new("", "", serde_json::json!({})));
|
||||
|
||||
events.push(StreamEvent::ToolCallDelta {
|
||||
tool_call: ToolCall {
|
||||
raw_arguments: Some(delta.to_string()),
|
||||
..current_tc
|
||||
},
|
||||
});
|
||||
}
|
||||
|
||||
/// Handle `response.output_item.done` for text and function call items.
|
||||
fn handle_output_item_done(
|
||||
state: &mut SseStreamState,
|
||||
json: &serde_json::Value,
|
||||
events: &mut Vec<StreamEvent>,
|
||||
) {
|
||||
let item_type = json
|
||||
.get("item")
|
||||
.and_then(|i| i.get("type"))
|
||||
.and_then(serde_json::Value::as_str);
|
||||
|
||||
match item_type {
|
||||
Some("message") => {
|
||||
if state.emitted_text_start {
|
||||
events.push(StreamEvent::TextEnd { text_id: None });
|
||||
state.emitted_text_start = false;
|
||||
}
|
||||
}
|
||||
Some("function_call") => {
|
||||
let item = json.get("item").unwrap_or(json);
|
||||
let call_id = item
|
||||
.get("call_id")
|
||||
.or_else(|| item.get("id"))
|
||||
.and_then(serde_json::Value::as_str)
|
||||
.unwrap_or("")
|
||||
.to_string();
|
||||
let name = item
|
||||
.get("name")
|
||||
.and_then(serde_json::Value::as_str)
|
||||
.unwrap_or("")
|
||||
.to_string();
|
||||
let args_str = item
|
||||
.get("arguments")
|
||||
.and_then(serde_json::Value::as_str)
|
||||
.unwrap_or("{}");
|
||||
let arguments =
|
||||
serde_json::from_str(args_str).unwrap_or_else(|_| serde_json::json!({}));
|
||||
|
||||
let mut tc = ToolCall::new(&call_id, &name, arguments);
|
||||
tc.raw_arguments = Some(args_str.to_string());
|
||||
|
||||
if let Some(existing) = state.tool_calls.iter_mut().find(|t| t.id == call_id) {
|
||||
existing.name.clone_from(&name);
|
||||
existing.arguments = tc.arguments.clone();
|
||||
existing.raw_arguments.clone_from(&tc.raw_arguments);
|
||||
} else {
|
||||
state.tool_calls.push(tc.clone());
|
||||
}
|
||||
|
||||
events.push(StreamEvent::ToolCallEnd { tool_call: tc });
|
||||
}
|
||||
_ => {}
|
||||
}
|
||||
}
|
||||
|
||||
/// Handle `response.completed` by extracting usage and building the final response.
|
||||
fn handle_response_completed(
|
||||
state: &mut SseStreamState,
|
||||
json: &serde_json::Value,
|
||||
events: &mut Vec<StreamEvent>,
|
||||
) {
|
||||
let response_data = json.get("response").unwrap_or(json);
|
||||
|
||||
if let Some(usage_data) = response_data.get("usage") {
|
||||
if let Ok(u) = serde_json::from_value::<ApiUsage>(usage_data.clone()) {
|
||||
state.usage = Usage {
|
||||
input_tokens: u.input_tokens,
|
||||
output_tokens: u.output_tokens,
|
||||
total_tokens: u.total_tokens.unwrap_or(u.input_tokens + u.output_tokens),
|
||||
reasoning_tokens: u
|
||||
.output_tokens_details
|
||||
.as_ref()
|
||||
.and_then(|d| d.reasoning_tokens),
|
||||
cache_read_tokens: u
|
||||
.input_tokens_details
|
||||
.as_ref()
|
||||
.and_then(|d| d.cached_tokens),
|
||||
..Usage::default()
|
||||
};
|
||||
}
|
||||
}
|
||||
|
||||
if let Some(id) = response_data
|
||||
.get("id")
|
||||
.and_then(serde_json::Value::as_str)
|
||||
{
|
||||
state.response_id = id.to_string();
|
||||
}
|
||||
if let Some(model) = response_data
|
||||
.get("model")
|
||||
.and_then(serde_json::Value::as_str)
|
||||
{
|
||||
state.response_model = model.to_string();
|
||||
}
|
||||
|
||||
let status = response_data
|
||||
.get("status")
|
||||
.and_then(serde_json::Value::as_str);
|
||||
let has_tool_calls = !state.tool_calls.is_empty();
|
||||
state.finish_reason = map_finish_reason(status, has_tool_calls);
|
||||
|
||||
state.raw_response = Some(response_data.clone());
|
||||
|
||||
let mut content_parts = Vec::new();
|
||||
if !state.accumulated_text.is_empty() {
|
||||
content_parts.push(ContentPart::text(&state.accumulated_text));
|
||||
}
|
||||
for tc in &state.tool_calls {
|
||||
content_parts.push(ContentPart::ToolCall(tc.clone()));
|
||||
}
|
||||
|
||||
let model = if state.response_model.is_empty() {
|
||||
state.model.clone()
|
||||
} else {
|
||||
state.response_model.clone()
|
||||
};
|
||||
|
||||
let response = Response {
|
||||
id: state.response_id.clone(),
|
||||
model,
|
||||
provider: "openai".to_string(),
|
||||
message: Message {
|
||||
role: Role::Assistant,
|
||||
content: content_parts,
|
||||
name: None,
|
||||
tool_call_id: None,
|
||||
},
|
||||
finish_reason: state.finish_reason.clone(),
|
||||
usage: state.usage.clone(),
|
||||
raw: state.raw_response.clone(),
|
||||
warnings: vec![],
|
||||
rate_limit: state.rate_limit.clone(),
|
||||
};
|
||||
|
||||
events.push(StreamEvent::finish(
|
||||
state.finish_reason.clone(),
|
||||
state.usage.clone(),
|
||||
response,
|
||||
));
|
||||
}
|
||||
|
||||
#[allow(clippy::unnecessary_literal_bound)]
|
||||
#[async_trait::async_trait]
|
||||
impl ProviderAdapter for Adapter {
|
||||
|
|
@ -98,36 +744,15 @@ impl ProviderAdapter for Adapter {
|
|||
}
|
||||
|
||||
async fn complete(&self, request: &Request) -> Result<Response, SdkError> {
|
||||
let api_messages: Vec<ApiMessage> = request
|
||||
.messages
|
||||
.iter()
|
||||
.map(|msg| {
|
||||
let role = match msg.role {
|
||||
Role::System | Role::Developer => "system",
|
||||
Role::User | Role::Tool => "user",
|
||||
Role::Assistant => "assistant",
|
||||
};
|
||||
ApiMessage {
|
||||
role: role.to_string(),
|
||||
content: msg.text(),
|
||||
}
|
||||
})
|
||||
.collect();
|
||||
let api_request = build_api_request(request, false);
|
||||
let url = format!("{}/responses", self.base_url);
|
||||
|
||||
let api_request = ApiRequest {
|
||||
model: request.model.clone(),
|
||||
messages: api_messages,
|
||||
temperature: request.temperature,
|
||||
max_tokens: request.max_tokens,
|
||||
top_p: request.top_p,
|
||||
stop: request.stop_sequences.clone(),
|
||||
};
|
||||
|
||||
let body = send_and_read_body(
|
||||
let (body, headers) = send_and_read_response(
|
||||
self.client
|
||||
.post("https://api.openai.com/v1/chat/completions")
|
||||
.post(&url)
|
||||
.bearer_auth(&self.api_key)
|
||||
.json(&api_request),
|
||||
.json(&api_request)
|
||||
.timeout(self.request_timeout),
|
||||
"openai",
|
||||
"type",
|
||||
)
|
||||
|
|
@ -138,41 +763,30 @@ impl ProviderAdapter for Adapter {
|
|||
message: format!("failed to parse OpenAI response: {e}"),
|
||||
})?;
|
||||
|
||||
let choice = api_resp.choices.first().ok_or_else(|| SdkError::Provider {
|
||||
kind: ProviderErrorKind::Server,
|
||||
detail: Box::new(ProviderErrorDetail::new("no choices in OpenAI response", "openai")),
|
||||
})?;
|
||||
let (content_parts, has_tool_calls) = parse_output(&api_resp.output);
|
||||
let finish_reason = map_finish_reason(api_resp.status.as_deref(), has_tool_calls);
|
||||
|
||||
let mut content_parts = Vec::new();
|
||||
if let Some(text) = &choice.message.content {
|
||||
if !text.is_empty() {
|
||||
content_parts.push(ContentPart::text(text));
|
||||
}
|
||||
}
|
||||
if let Some(tool_calls) = &choice.message.tool_calls {
|
||||
for tc in tool_calls {
|
||||
let arguments = serde_json::from_str(&tc.function.arguments)
|
||||
.unwrap_or_else(|_| serde_json::json!({}));
|
||||
content_parts.push(ContentPart::ToolCall(ToolCall::new(
|
||||
&tc.id,
|
||||
&tc.function.name,
|
||||
arguments,
|
||||
)));
|
||||
}
|
||||
}
|
||||
|
||||
let finish_reason = map_finish_reason(choice.finish_reason.as_deref());
|
||||
|
||||
let usage = api_resp.usage.as_ref().map_or_else(Usage::default, |u| Usage {
|
||||
input_tokens: u.prompt_tokens,
|
||||
output_tokens: u.completion_tokens,
|
||||
total_tokens: u.total_tokens,
|
||||
..Usage::default()
|
||||
});
|
||||
let usage = api_resp
|
||||
.usage
|
||||
.as_ref()
|
||||
.map_or_else(Usage::default, |u| Usage {
|
||||
input_tokens: u.input_tokens,
|
||||
output_tokens: u.output_tokens,
|
||||
total_tokens: u.total_tokens.unwrap_or(u.input_tokens + u.output_tokens),
|
||||
reasoning_tokens: u
|
||||
.output_tokens_details
|
||||
.as_ref()
|
||||
.and_then(|d| d.reasoning_tokens),
|
||||
cache_read_tokens: u
|
||||
.input_tokens_details
|
||||
.as_ref()
|
||||
.and_then(|d| d.cached_tokens),
|
||||
..Usage::default()
|
||||
});
|
||||
|
||||
Ok(Response {
|
||||
id: api_resp.id,
|
||||
model: api_resp.model,
|
||||
model: api_resp.model.unwrap_or_else(|| request.model.clone()),
|
||||
provider: "openai".to_string(),
|
||||
message: Message {
|
||||
role: Role::Assistant,
|
||||
|
|
@ -184,13 +798,73 @@ impl ProviderAdapter for Adapter {
|
|||
usage,
|
||||
raw: serde_json::from_str(&body).ok(),
|
||||
warnings: vec![],
|
||||
rate_limit: None,
|
||||
rate_limit: parse_rate_limit_headers(&headers),
|
||||
})
|
||||
}
|
||||
|
||||
async fn stream(&self, _request: &Request) -> Result<StreamEventStream, SdkError> {
|
||||
Err(SdkError::Configuration {
|
||||
message: "streaming not yet implemented".to_string(),
|
||||
async fn stream(&self, request: &Request) -> Result<StreamEventStream, SdkError> {
|
||||
let api_request = build_api_request(request, true);
|
||||
let url = format!("{}/responses", self.base_url);
|
||||
|
||||
let http_resp = self
|
||||
.client
|
||||
.post(&url)
|
||||
.bearer_auth(&self.api_key)
|
||||
.json(&api_request)
|
||||
.send()
|
||||
.await
|
||||
.map_err(|e| SdkError::Network {
|
||||
message: e.to_string(),
|
||||
})?;
|
||||
|
||||
let status = http_resp.status();
|
||||
if !status.is_success() {
|
||||
let retry_after = parse_retry_after(http_resp.headers());
|
||||
let body = http_resp.text().await.map_err(|e| SdkError::Network {
|
||||
message: e.to_string(),
|
||||
})?;
|
||||
let (msg, code, raw) = parse_error_body(&body, "type");
|
||||
return Err(crate::error::error_from_status_code(
|
||||
status.as_u16(),
|
||||
msg,
|
||||
"openai".to_string(),
|
||||
code,
|
||||
raw,
|
||||
retry_after,
|
||||
));
|
||||
}
|
||||
|
||||
let model = request.model.clone();
|
||||
let rate_limit = parse_rate_limit_headers(http_resp.headers());
|
||||
let byte_stream = http_resp.bytes_stream();
|
||||
|
||||
let state = SseStreamState {
|
||||
byte_stream: Box::pin(byte_stream),
|
||||
buffer: String::new(),
|
||||
model,
|
||||
response_id: String::new(),
|
||||
response_model: String::new(),
|
||||
accumulated_text: String::new(),
|
||||
tool_calls: Vec::new(),
|
||||
usage: Usage::default(),
|
||||
finish_reason: FinishReason::Stop,
|
||||
emitted_start: false,
|
||||
emitted_text_start: false,
|
||||
raw_response: None,
|
||||
rate_limit,
|
||||
};
|
||||
|
||||
let stream = futures::stream::unfold(state, |mut state| async move {
|
||||
let events = process_next_sse_events(&mut state).await;
|
||||
let items: Vec<Result<StreamEvent, SdkError>> = match events {
|
||||
Ok(events) if events.is_empty() => return None,
|
||||
Ok(events) => events.into_iter().map(Ok).collect(),
|
||||
Err(e) => vec![Err(e)],
|
||||
};
|
||||
Some((futures::stream::iter(items), state))
|
||||
})
|
||||
.flatten();
|
||||
|
||||
Ok(Box::pin(stream))
|
||||
}
|
||||
}
|
||||
|
|
|
|||
1114
crates/unified-llm/src/providers/openai_compatible.rs
Normal file
1114
crates/unified-llm/src/providers/openai_compatible.rs
Normal file
File diff suppressed because it is too large
Load diff
|
|
@ -27,16 +27,21 @@ where
|
|||
}
|
||||
|
||||
// Check Retry-After
|
||||
if let Some(retry_after) = err.retry_after() {
|
||||
let delay = if let Some(retry_after) = err.retry_after() {
|
||||
if retry_after > policy.max_delay {
|
||||
return Err(err);
|
||||
}
|
||||
tokio::time::sleep(std::time::Duration::from_secs_f64(retry_after)).await;
|
||||
retry_after
|
||||
} else {
|
||||
let delay = policy.delay_for_attempt(attempt);
|
||||
tokio::time::sleep(std::time::Duration::from_secs_f64(delay)).await;
|
||||
policy.delay_for_attempt(attempt)
|
||||
};
|
||||
|
||||
if let Some(ref on_retry) = policy.on_retry {
|
||||
on_retry(&err, attempt, delay);
|
||||
}
|
||||
|
||||
tokio::time::sleep(std::time::Duration::from_secs_f64(delay)).await;
|
||||
|
||||
attempt += 1;
|
||||
}
|
||||
}
|
||||
|
|
@ -46,6 +51,7 @@ where
|
|||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use crate::types::RetryPolicy;
|
||||
use std::sync::atomic::{AtomicU32, Ordering};
|
||||
use std::sync::Arc;
|
||||
|
||||
|
|
@ -243,4 +249,47 @@ mod tests {
|
|||
// Should have waited ~0.01s, not ~10s
|
||||
assert!(elapsed.as_secs_f64() < 1.0);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn retry_invokes_on_retry_callback() {
|
||||
let retry_attempts = Arc::new(AtomicU32::new(0));
|
||||
let retry_attempts_clone = retry_attempts.clone();
|
||||
|
||||
let policy = RetryPolicy {
|
||||
max_retries: 2,
|
||||
base_delay: 0.001,
|
||||
jitter: false,
|
||||
on_retry: Some(Arc::new(move |_err, _attempt, _delay| {
|
||||
retry_attempts_clone.fetch_add(1, Ordering::SeqCst);
|
||||
})),
|
||||
..Default::default()
|
||||
};
|
||||
|
||||
let call_count = Arc::new(AtomicU32::new(0));
|
||||
let cc = call_count.clone();
|
||||
|
||||
let result = retry(&policy, || {
|
||||
let cc = cc.clone();
|
||||
async move {
|
||||
let count = cc.fetch_add(1, Ordering::SeqCst);
|
||||
if count < 2 {
|
||||
Err(SdkError::Provider {
|
||||
kind: crate::error::ProviderErrorKind::Server,
|
||||
detail: Box::new(crate::error::ProviderErrorDetail {
|
||||
status_code: Some(500),
|
||||
..crate::error::ProviderErrorDetail::new("error", "test")
|
||||
}),
|
||||
})
|
||||
} else {
|
||||
Ok(99)
|
||||
}
|
||||
}
|
||||
})
|
||||
.await;
|
||||
|
||||
assert_eq!(result.unwrap(), 99);
|
||||
assert_eq!(call_count.load(Ordering::SeqCst), 3);
|
||||
// on_retry should have been called twice (before each retry)
|
||||
assert_eq!(retry_attempts.load(Ordering::SeqCst), 2);
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -20,8 +20,15 @@ pub struct Tool {
|
|||
|
||||
impl Tool {
|
||||
/// Create a passive tool (no execute handler).
|
||||
#[must_use]
|
||||
///
|
||||
/// # Panics
|
||||
///
|
||||
/// Panics if the tool name is invalid (see [`validate_tool_name`]).
|
||||
#[must_use]
|
||||
pub fn passive(name: &str, description: &str, parameters: serde_json::Value) -> Self {
|
||||
if let Err(e) = validate_tool_name(name) {
|
||||
panic!("Invalid tool name: {e}");
|
||||
}
|
||||
Self {
|
||||
definition: ToolDefinition {
|
||||
name: name.to_string(),
|
||||
|
|
@ -33,6 +40,10 @@ impl Tool {
|
|||
}
|
||||
|
||||
/// Create an active tool with an execute handler.
|
||||
///
|
||||
/// # Panics
|
||||
///
|
||||
/// Panics if the tool name is invalid (see [`validate_tool_name`]).
|
||||
pub fn active<F, Fut>(
|
||||
name: &str,
|
||||
description: &str,
|
||||
|
|
@ -43,6 +54,9 @@ impl Tool {
|
|||
F: Fn(serde_json::Value) -> Fut + Send + Sync + 'static,
|
||||
Fut: Future<Output = Result<serde_json::Value, String>> + Send + 'static,
|
||||
{
|
||||
if let Err(e) = validate_tool_name(name) {
|
||||
panic!("Invalid tool name: {e}");
|
||||
}
|
||||
Self {
|
||||
definition: ToolDefinition {
|
||||
name: name.to_string(),
|
||||
|
|
@ -121,11 +135,15 @@ pub async fn execute_all_tools(
|
|||
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,
|
||||
},
|
||||
}
|
||||
}
|
||||
|
|
@ -135,6 +153,8 @@ pub async fn execute_all_tools(
|
|||
"Unknown tool: {call_name}"
|
||||
)),
|
||||
is_error: true,
|
||||
image_data: None,
|
||||
image_media_type: None,
|
||||
},
|
||||
}
|
||||
}
|
||||
|
|
@ -340,4 +360,25 @@ mod tests {
|
|||
assert!(!results[0].is_error);
|
||||
assert!(results[1].is_error);
|
||||
}
|
||||
|
||||
#[test]
|
||||
#[should_panic(expected = "Invalid tool name")]
|
||||
fn passive_tool_panics_on_invalid_name() {
|
||||
let _ = Tool::passive(
|
||||
"1invalid",
|
||||
"bad name",
|
||||
serde_json::json!({"type": "object"}),
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
#[should_panic(expected = "Invalid tool name")]
|
||||
fn active_tool_panics_on_invalid_name() {
|
||||
Tool::active(
|
||||
"my-tool",
|
||||
"bad name",
|
||||
serde_json::json!({"type": "object"}),
|
||||
|_args| async { Ok(serde_json::json!("result")) },
|
||||
);
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,5 +1,7 @@
|
|||
use crate::error::SdkError;
|
||||
use serde::{Deserialize, Serialize};
|
||||
use std::collections::HashMap;
|
||||
use std::sync::Arc;
|
||||
|
||||
// --- 3.2 Role ---
|
||||
|
||||
|
|
@ -75,6 +77,10 @@ pub struct ToolResult {
|
|||
pub tool_call_id: String,
|
||||
pub content: serde_json::Value,
|
||||
pub is_error: bool,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub image_data: Option<Vec<u8>>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub image_media_type: Option<String>,
|
||||
}
|
||||
|
||||
// --- 3.3 ContentPart ---
|
||||
|
|
@ -149,6 +155,8 @@ impl Message {
|
|||
tool_call_id: id.clone(),
|
||||
content: serde_json::Value::String(content.into()),
|
||||
is_error,
|
||||
image_data: None,
|
||||
image_media_type: None,
|
||||
})],
|
||||
name: None,
|
||||
tool_call_id: Some(id),
|
||||
|
|
@ -503,13 +511,31 @@ impl Default for AdapterTimeout {
|
|||
|
||||
// --- 6.6 RetryPolicy ---
|
||||
|
||||
#[derive(Debug, Clone)]
|
||||
/// Callback invoked before each retry attempt with (error, attempt, delay in seconds).
|
||||
pub type OnRetryCallback = Arc<dyn Fn(&SdkError, u32, f64) + Send + Sync>;
|
||||
|
||||
#[derive(Clone)]
|
||||
pub struct RetryPolicy {
|
||||
pub max_retries: u32,
|
||||
pub base_delay: f64,
|
||||
pub max_delay: f64,
|
||||
pub backoff_multiplier: f64,
|
||||
pub jitter: bool,
|
||||
/// Called before each retry with (error, attempt number, delay in seconds).
|
||||
pub on_retry: Option<OnRetryCallback>,
|
||||
}
|
||||
|
||||
impl std::fmt::Debug for RetryPolicy {
|
||||
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
f.debug_struct("RetryPolicy")
|
||||
.field("max_retries", &self.max_retries)
|
||||
.field("base_delay", &self.base_delay)
|
||||
.field("max_delay", &self.max_delay)
|
||||
.field("backoff_multiplier", &self.backoff_multiplier)
|
||||
.field("jitter", &self.jitter)
|
||||
.field("on_retry", &self.on_retry.as_ref().map(|_| "..."))
|
||||
.finish()
|
||||
}
|
||||
}
|
||||
|
||||
impl Default for RetryPolicy {
|
||||
|
|
@ -520,6 +546,7 @@ impl Default for RetryPolicy {
|
|||
max_delay: 60.0,
|
||||
backoff_multiplier: 2.0,
|
||||
jitter: true,
|
||||
on_retry: None,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
@ -540,6 +567,22 @@ impl RetryPolicy {
|
|||
}
|
||||
}
|
||||
|
||||
// --- 4.6 ObjectStreamEvent ---
|
||||
|
||||
/// Events yielded by `stream_object()` for streaming structured output.
|
||||
#[derive(Debug, Clone)]
|
||||
pub enum ObjectStreamEvent {
|
||||
/// A new partial parse of the accumulated JSON text.
|
||||
Partial { object: serde_json::Value },
|
||||
/// A raw stream event from the underlying provider stream.
|
||||
Delta { event: StreamEvent },
|
||||
/// The stream completed with a fully parsed object and response.
|
||||
Complete {
|
||||
object: serde_json::Value,
|
||||
response: Box<Response>,
|
||||
},
|
||||
}
|
||||
|
||||
// --- 4.3 GenerateResult / StepResult ---
|
||||
|
||||
#[derive(Debug, Clone)]
|
||||
|
|
@ -905,6 +948,7 @@ mod tests {
|
|||
max_delay: 60.0,
|
||||
backoff_multiplier: 2.0,
|
||||
jitter: false,
|
||||
..Default::default()
|
||||
};
|
||||
assert!((policy.delay_for_attempt(0) - 1.0).abs() < f64::EPSILON);
|
||||
assert!((policy.delay_for_attempt(1) - 2.0).abs() < f64::EPSILON);
|
||||
|
|
@ -920,6 +964,7 @@ mod tests {
|
|||
max_delay: 5.0,
|
||||
backoff_multiplier: 2.0,
|
||||
jitter: false,
|
||||
..Default::default()
|
||||
};
|
||||
assert!((policy.delay_for_attempt(5) - 5.0).abs() < f64::EPSILON);
|
||||
}
|
||||
|
|
@ -932,6 +977,7 @@ mod tests {
|
|||
max_delay: 60.0,
|
||||
backoff_multiplier: 2.0,
|
||||
jitter: true,
|
||||
..Default::default()
|
||||
};
|
||||
let delay = policy.delay_for_attempt(0);
|
||||
// base * 0.5 to base * 1.5 => 0.5 to 1.5
|
||||
|
|
@ -972,6 +1018,19 @@ mod tests {
|
|||
assert_eq!(deserialized, tc);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn tool_result_with_image_data() {
|
||||
let result = ToolResult {
|
||||
tool_call_id: "call_1".into(),
|
||||
content: serde_json::json!("screenshot taken"),
|
||||
is_error: false,
|
||||
image_data: Some(vec![0x89, 0x50, 0x4E, 0x47]),
|
||||
image_media_type: Some("image/png".into()),
|
||||
};
|
||||
assert!(result.image_data.is_some());
|
||||
assert_eq!(result.image_media_type.as_deref(), Some("image/png"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn tool_call_new_constructor() {
|
||||
let tc = ToolCall::new("c1", "test", serde_json::json!({}));
|
||||
|
|
|
|||
|
|
@ -510,10 +510,10 @@ RECORD DocumentData:
|
|||
|
||||
```
|
||||
RECORD ToolCallData:
|
||||
id : String -- unique identifier for this call (provider-assigned)
|
||||
name : String -- tool name
|
||||
arguments : Dict | String -- parsed JSON arguments or raw argument string
|
||||
type : String -- "function" (default) or "custom"
|
||||
id : String -- unique identifier for this call (provider-assigned)
|
||||
name : String -- tool name
|
||||
arguments : Dict -- parsed JSON arguments
|
||||
raw_arguments : String | None -- raw argument string before parsing (for debugging)
|
||||
```
|
||||
|
||||
The `id` field is assigned by the provider and is required for linking tool results back to calls. For providers that do not assign unique IDs (e.g., Gemini), the adapter must generate synthetic unique IDs (e.g., `"call_" + random_uuid()`) and maintain a mapping to the function name.
|
||||
|
|
@ -617,6 +617,8 @@ RECORD FinishReason:
|
|||
raw : String | None -- the provider's native finish reason string
|
||||
```
|
||||
|
||||
**Note for statically-typed languages:** An enum (discriminated union) with variants `Stop`, `Length`, `ToolCalls`, `ContentFilter`, `Error`, and `Other(String)` is an acceptable representation. The provider's raw finish reason string is available in `Response.raw` (the full provider response). A separate `raw` field on FinishReason itself is optional in typed implementations.
|
||||
|
||||
Unified reason values:
|
||||
|
||||
| Value | Meaning |
|
||||
|
|
@ -632,10 +634,14 @@ Provider finish reason mapping:
|
|||
|
||||
| Provider | Provider Value | Unified Value |
|
||||
|-----------|-------------------|------------------|
|
||||
| OpenAI | stop | stop |
|
||||
| OpenAI | length | length |
|
||||
| OpenAI | tool_calls | tool_calls |
|
||||
| OpenAI | content_filter | content_filter |
|
||||
| OpenAI (Responses API) | completed | stop |
|
||||
| OpenAI (Responses API) | incomplete | length |
|
||||
| OpenAI (Responses API) | failed | error |
|
||||
| OpenAI (Responses API) | (has function_call items) | tool_calls |
|
||||
| OpenAI (Chat Completions) | stop | stop |
|
||||
| OpenAI (Chat Completions) | length | length |
|
||||
| OpenAI (Chat Completions) | tool_calls | tool_calls |
|
||||
| OpenAI (Chat Completions) | content_filter | content_filter |
|
||||
| Anthropic | end_turn | stop |
|
||||
| Anthropic | stop_sequence | stop |
|
||||
| Anthropic | max_tokens | length |
|
||||
|
|
@ -646,7 +652,7 @@ Provider finish reason mapping:
|
|||
| Gemini | RECITATION | content_filter |
|
||||
| Gemini | (has tool calls) | tool_calls |
|
||||
|
||||
Note: Gemini does not have a dedicated "tool_calls" finish reason. The adapter infers it from the presence of `functionCall` parts in the response.
|
||||
Note: Gemini does not have a dedicated "tool_calls" finish reason. The adapter infers it from the presence of `functionCall` parts in the response. Similarly, the OpenAI Responses API uses `completed`/`incomplete`/`failed` status rather than Chat Completions-style finish reasons, and tool calls are inferred from the presence of `function_call` output items.
|
||||
|
||||
### 3.9 Usage
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue