Implement spec gaps: abort signal, StreamResult, provider_options, max_tool_rounds fix

- Add abort signal support using CancellationToken for cooperative cancellation
  of generate() and stream() calls
- Fix max_tool_rounds=0 to skip tool execution entirely (was executing first round)
- Add StreamResult wrapper with response(), text_stream(), partial_response()
  and multi-step tool loop support in high-level stream()
- Add OpenAI metadata and provider_options.openai pass-through to Responses API
- Add Gemini provider_options.gemini pass-through (safety settings, cached content)
- Add OpenAI-compatible provider_options.<name> pass-through using adapter name

Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
This commit is contained in:
Bryan Helmkamp 2026-02-20 11:23:22 -04:00
parent ed07d43335
commit 5f2385b87b
8 changed files with 1134 additions and 23 deletions

1
Cargo.lock generated
View file

@ -1401,6 +1401,7 @@ dependencies = [
"thiserror",
"tokio",
"tokio-stream",
"tokio-util",
"uuid",
]

View file

@ -29,3 +29,4 @@ tokio-stream = "0.1"
async-trait = "0.1"
base64 = "0.22"
bytes = "1"
tokio-util = "0.7"

View file

@ -23,6 +23,7 @@ async-trait.workspace = true
reqwest.workspace = true
base64.workspace = true
bytes.workspace = true
tokio-util.workspace = true
[dev-dependencies]
dotenvy.workspace = true

View file

@ -8,10 +8,12 @@ use crate::types::{
ResponseFormatType, RetryPolicy, StepResult, StreamEvent, TimeoutConfig, ToolCall, ToolChoice,
ToolDefinition, Usage,
};
use futures::StreamExt;
use futures::{Stream, StreamExt};
use std::pin::Pin;
use std::sync::Arc;
use std::task::{Context, Poll};
use tokio::sync::OnceCell;
use tokio_util::sync::CancellationToken;
/// Module-level default client (Section 2.5).
static DEFAULT_CLIENT: OnceCell<Arc<Client>> = OnceCell::const_new();
@ -95,6 +97,7 @@ fn build_generate_result(steps: Vec<StepResult>, total_usage: Usage) -> Generate
/// # Panics
///
/// Panics if a tool's `execute` handler is `None` when matched during tool execution.
#[allow(clippy::too_many_lines)]
pub async fn generate(params: GenerateParams) -> Result<GenerateResult, SdkError> {
let client = params.client.clone().unwrap_or_else(get_default_client);
let retry_policy = RetryPolicy {
@ -112,12 +115,22 @@ pub async fn generate(params: GenerateParams) -> Result<GenerateResult, SdkError
let max_tool_rounds = params.max_tool_rounds;
let abort_signal = params.abort_signal.clone();
let generate_future = async {
let mut steps: Vec<StepResult> = Vec::new();
let mut total_usage = Usage::default();
let mut round = 0u32;
loop {
if let Some(ref token) = abort_signal {
if token.is_cancelled() {
return Err(SdkError::Abort {
message: "Generation aborted by cancellation token".into(),
});
}
}
let request = build_request(&params, &messages, tool_definitions.as_deref());
let client_ref = client.clone();
@ -148,6 +161,7 @@ pub async fn generate(params: GenerateParams) -> Result<GenerateResult, SdkError
if !tool_calls.is_empty()
&& response.finish_reason == FinishReason::ToolCalls
&& params.tools.is_some()
&& max_tool_rounds > 0
{
let tools = params.tools.as_ref().expect("checked above");
if tools.iter().any(|t| t.is_active()) {
@ -175,6 +189,14 @@ pub async fn generate(params: GenerateParams) -> Result<GenerateResult, SdkError
break;
}
if let Some(ref token) = abort_signal {
if token.is_cancelled() {
return Err(SdkError::Abort {
message: "Generation aborted by cancellation token".into(),
});
}
}
let last = steps.last().expect("just pushed");
messages.push(last.response.message.clone());
for result in &last.tool_results {
@ -228,6 +250,8 @@ pub struct GenerateParams {
pub max_retries: u32,
pub timeout: Option<TimeoutConfig>,
pub client: Option<Arc<Client>>,
/// Cancellation token to abort generation (Section 4.8).
pub abort_signal: Option<CancellationToken>,
/// Custom stop condition checked after each tool round (Section 4.3).
pub stop_when: Option<StopCondition>,
}
@ -254,6 +278,7 @@ impl GenerateParams {
max_retries: 2,
timeout: None,
client: None,
abort_signal: None,
stop_when: None,
}
}
@ -366,6 +391,12 @@ impl GenerateParams {
self
}
#[must_use]
pub fn abort_signal(mut self, token: CancellationToken) -> Self {
self.abort_signal = Some(token);
self
}
/// Set a custom stop condition for the tool loop (Section 4.3).
///
/// The callback receives the accumulated steps so far and returns `true`
@ -453,15 +484,237 @@ impl Default for StreamAccumulator {
}
}
/// Wraps a streaming response with an internal `StreamAccumulator` and convenience methods.
///
/// Implements `Stream<Item = Result<StreamEvent, SdkError>>` so it can be used
/// as a drop-in replacement for `StreamEventStream`. Also supports multi-step
/// tool loops when active tools are provided.
pub struct StreamResult {
inner: StreamEventStream,
accumulator: StreamAccumulator,
}
impl StreamResult {
fn new(inner: StreamEventStream) -> Self {
Self {
inner,
accumulator: StreamAccumulator::new(),
}
}
/// Returns the accumulated response after the stream has ended.
#[must_use]
pub const fn response(&self) -> Option<&Response> {
self.accumulator.response()
}
/// Returns the current partially accumulated response state.
#[must_use]
pub const fn partial_response(&self) -> Option<&Response> {
self.accumulator.response()
}
/// Returns a stream that yields only text delta strings.
#[must_use]
pub fn text_stream(self) -> Pin<Box<dyn Stream<Item = Result<String, SdkError>> + Send>> {
Box::pin(self.filter_map(|result| {
futures::future::ready(match result {
Ok(StreamEvent::TextDelta { delta, .. }) => Some(Ok(delta)),
Err(e) => Some(Err(e)),
_ => None,
})
}))
}
}
impl Stream for StreamResult {
type Item = Result<StreamEvent, SdkError>;
fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
let inner = self.inner.as_mut();
match inner.poll_next(cx) {
Poll::Ready(Some(Ok(event))) => {
self.accumulator.process(&event);
Poll::Ready(Some(Ok(event)))
}
other => other,
}
}
}
/// High-level streaming generation (Section 4.4).
/// Returns a `StreamEventStream` that the caller can iterate over.
/// Returns a `StreamResult` that the caller can iterate over.
/// Supports multi-step tool loops when active tools are provided.
///
/// # Errors
///
/// 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
pub async fn stream(params: GenerateParams) -> Result<StreamResult, SdkError> {
let inner = stream_with_tool_loop(params).await?;
Ok(StreamResult::new(inner))
}
/// Streaming generation with multi-step tool loop support.
///
/// When active tools are provided and the model returns tool calls:
/// - Collects the stream to get the complete first response
/// - Executes tools concurrently
/// - Starts a new stream with updated conversation
/// - Yields all events from all rounds seamlessly
/// - Continues until no more tool calls or `max_tool_rounds` reached
///
/// # Errors
///
/// Returns `SdkError::Configuration` if both `prompt` and `messages` are set,
/// or any provider error encountered during streaming setup.
#[allow(clippy::too_many_lines)]
async fn stream_with_tool_loop(params: GenerateParams) -> Result<StreamEventStream, SdkError> {
let client = params.client.clone().unwrap_or_else(get_default_client);
let mut messages = build_initial_messages(&params)?;
let tool_definitions: Option<Vec<ToolDefinition>> = params
.tools
.as_ref()
.map(|tools| tools.iter().map(|t| t.definition.clone()).collect());
let abort_signal = params.abort_signal.clone();
let max_tool_rounds = params.max_tool_rounds;
let has_active_tools = max_tool_rounds > 0
&& params
.tools
.as_ref()
.is_some_and(|tools| tools.iter().any(|t| t.is_active()));
if !has_active_tools {
// No tool loop needed, just stream directly
return stream_generate_raw(&client, &params, &messages, tool_definitions.as_deref())
.await;
}
// Tool loop: collect events from each round, execute tools, continue
let (tx, rx) = tokio::sync::mpsc::channel::<Result<StreamEvent, SdkError>>(64);
let tools = params.tools.clone();
tokio::spawn(async move {
let mut round = 0u32;
loop {
if let Some(ref token) = abort_signal {
if token.is_cancelled() {
let _ = tx
.send(Err(SdkError::Abort {
message: "Stream aborted by cancellation token".into(),
}))
.await;
return;
}
}
let request = build_request(&params, &messages, tool_definitions.as_deref());
let stream_result = client.stream(&request).await;
let mut inner_stream = match stream_result {
Ok(s) => s,
Err(e) => {
let _ = tx.send(Err(e)).await;
return;
}
};
// Collect stream and forward events, accumulating for tool call detection
let mut accumulator = StreamAccumulator::new();
while let Some(item) = inner_stream.next().await {
if let Some(ref token) = abort_signal {
if token.is_cancelled() {
let _ = tx
.send(Err(SdkError::Abort {
message: "Stream aborted by cancellation token".into(),
}))
.await;
return;
}
}
if let Ok(event) = &item {
accumulator.process(event);
} else {
let _ = tx.send(item).await;
return;
}
// Forward the event to the consumer
if tx.send(item).await.is_err() {
return; // Consumer dropped
}
}
// Check if we should continue with tool calls
let response = match accumulator.response() {
Some(r) => r.clone(),
None => return, // No response accumulated, stream ended
};
let tool_calls = response.tool_calls();
if tool_calls.is_empty()
|| response.finish_reason != FinishReason::ToolCalls
|| round >= max_tool_rounds
{
return; // No more tool rounds needed
}
// Execute tools
let Some(tool_list) = &tools else { return };
let tool_refs: Vec<&Tool> = tool_list.iter().map(std::convert::AsRef::as_ref).collect();
let tool_results = execute_all_tools(&tool_refs, &tool_calls).await;
if tool_results.is_empty() {
return;
}
// Append assistant message and tool results to conversation
messages.push(response.message.clone());
for result in &tool_results {
messages.push(Message::tool_result(
&result.tool_call_id,
result.content.to_string(),
result.is_error,
));
}
round += 1;
}
});
Ok(Box::pin(tokio_stream::wrappers::ReceiverStream::new(rx)))
}
/// Internal single-round streaming (no tool loop). Used by `stream_object()`.
async fn stream_generate_raw(
client: &Arc<Client>,
params: &GenerateParams,
messages: &[Message],
tool_definitions: Option<&[ToolDefinition]>,
) -> Result<StreamEventStream, SdkError> {
let request = build_request(params, messages, tool_definitions);
let inner_stream = client.stream(&request).await?;
if let Some(ref token) = params.abort_signal {
let token = token.clone();
let mapped = inner_stream.map(move |item| {
if token.is_cancelled() {
return Err(SdkError::Abort {
message: "Stream aborted by cancellation token".into(),
});
}
item
});
Ok(Box::pin(mapped))
} else {
Ok(inner_stream)
}
}
/// High-level streaming generation (Section 4.4).
@ -481,8 +734,7 @@ pub async fn stream_generate(params: GenerateParams) -> Result<StreamEventStream
.as_ref()
.map(|tools| tools.iter().map(|t| t.definition.clone()).collect());
let request = build_request(&params, &messages, tool_definitions.as_deref());
client.stream(&request).await
stream_generate_raw(&client, &params, &messages, tool_definitions.as_deref()).await
}
/// Structured output generation with schema validation (Section 4.5).
@ -1294,4 +1546,414 @@ mod tests {
let has_error = results.iter().any(|r| r.is_err());
assert!(has_error, "Expected an error for invalid final JSON");
}
#[tokio::test]
async fn generate_abort_signal_before_call() {
let token = CancellationToken::new();
token.cancel();
let result = generate(
GenerateParams::new("mock-model")
.prompt("Hello")
.client(mock_client("Hi"))
.abort_signal(token),
)
.await;
assert!(result.is_err());
assert!(matches!(result.unwrap_err(), SdkError::Abort { .. }));
}
#[tokio::test]
async fn generate_abort_signal_between_tool_rounds() {
let call_count = Arc::new(AtomicU32::new(0));
let token = CancellationToken::new();
let token_clone = token.clone();
// Provider that always returns tool calls
struct AlwaysToolCallProvider {
call_count: Arc<AtomicU32>,
cancel_token: CancellationToken,
}
#[async_trait::async_trait]
impl ProviderAdapter for AlwaysToolCallProvider {
fn name(&self) -> &str {
"mock"
}
async fn complete(&self, _request: &Request) -> Result<Response, SdkError> {
let count = self.call_count.fetch_add(1, Ordering::SeqCst);
// Cancel after first call completes
if count == 0 {
self.cancel_token.cancel();
}
Ok(Response {
id: format!("resp_{count}"),
model: "mock-model".into(),
provider: "mock".into(),
message: Message {
role: Role::Assistant,
content: vec![ContentPart::ToolCall(ToolCall::new(
format!("call_{count}"),
"get_weather",
serde_json::json!({"city": "SF"}),
))],
name: None,
tool_call_id: None,
},
finish_reason: FinishReason::ToolCalls,
usage: Usage::default(),
raw: None,
warnings: vec![],
rate_limit: None,
})
}
async fn stream(
&self,
_request: &Request,
) -> Result<StreamEventStream, SdkError> {
Ok(Box::pin(stream::empty()))
}
}
let provider: Arc<dyn ProviderAdapter> = Arc::new(AlwaysToolCallProvider {
call_count: call_count.clone(),
cancel_token: token_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?")
.tools(vec![Tool::active(
"get_weather",
"Get weather",
serde_json::json!({"type": "object", "properties": {"city": {"type": "string"}}}),
|_args| async { Ok(serde_json::json!("72F")) },
)])
.max_tool_rounds(10)
.abort_signal(token)
.client(client),
)
.await;
assert!(result.is_err());
assert!(matches!(result.unwrap_err(), SdkError::Abort { .. }));
// Should have only made 1 call before aborting
assert_eq!(call_count.load(Ordering::SeqCst), 1);
}
#[tokio::test]
async fn stream_abort_signal_terminates_stream() {
let token = CancellationToken::new();
let token_clone = token.clone();
// Create a mock that produces events, but cancel after stream starts
let client = mock_client("Hello stream!");
token_clone.cancel();
let mut stream_result = stream(
GenerateParams::new("mock-model")
.prompt("Hi")
.client(client)
.abort_signal(token),
)
.await
.unwrap();
let first = stream_result.next().await.unwrap();
assert!(first.is_err());
assert!(matches!(first.unwrap_err(), SdkError::Abort { .. }));
}
#[tokio::test]
async fn generate_max_tool_rounds_zero_skips_tool_execution() {
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 tool_executed = Arc::new(AtomicU32::new(0));
let tool_executed_clone = tool_executed.clone();
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"}}}),
move |_args| {
let counter = tool_executed_clone.clone();
async move {
counter.fetch_add(1, Ordering::SeqCst);
Ok(serde_json::json!("72F"))
}
},
)])
.max_tool_rounds(0)
.client(client),
)
.await
.unwrap();
// Should return after first LLM call without executing any tools
assert_eq!(result.steps.len(), 1);
assert_eq!(call_count.load(Ordering::SeqCst), 1);
assert_eq!(tool_executed.load(Ordering::SeqCst), 0);
// The tool results should be empty since tools were not executed
assert!(result.tool_results.is_empty());
}
#[test]
fn generate_params_abort_signal_builder() {
let token = CancellationToken::new();
let params = GenerateParams::new("test-model")
.abort_signal(token);
assert!(params.abort_signal.is_some());
}
#[tokio::test]
async fn stream_result_accumulates_response() {
let client = mock_client("Hello!");
let mut result = stream(
GenerateParams::new("mock-model")
.prompt("Hi")
.client(client),
)
.await
.unwrap();
assert!(result.response().is_none());
assert!(result.partial_response().is_none());
// Consume all events
while result.next().await.is_some() {}
assert!(result.response().is_some());
assert_eq!(result.response().unwrap().text(), "Hello!");
}
#[tokio::test]
async fn stream_result_text_stream() {
let client = streaming_json_mock_client(vec!["Hello", " ", "world"]);
let result = stream(
GenerateParams::new("mock-model")
.prompt("Hi")
.client(client),
)
.await
.unwrap();
let texts: Vec<String> = result
.text_stream()
.filter_map(|r| futures::future::ready(r.ok()))
.collect()
.await;
assert_eq!(texts, vec!["Hello", " ", "world"]);
}
/// Mock provider that streams tool calls then text on second stream
struct StreamingToolCallMockProvider {
call_count: Arc<AtomicU32>,
}
#[async_trait::async_trait]
impl ProviderAdapter for StreamingToolCallMockProvider {
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("fallback"),
finish_reason: FinishReason::Stop,
usage: Usage::default(),
raw: None,
warnings: vec![],
rate_limit: None,
})
}
async fn stream(
&self,
_request: &Request,
) -> Result<StreamEventStream, SdkError> {
let count = self.call_count.fetch_add(1, Ordering::SeqCst);
if count == 0 {
// First stream: return tool call
let tool_call = ToolCall::new("call_1", "get_weather", serde_json::json!({"city": "SF"}));
let response = Response {
id: "resp_1".into(),
model: "mock-model".into(),
provider: "mock".into(),
message: Message {
role: Role::Assistant,
content: vec![ContentPart::ToolCall(tool_call.clone())],
name: None,
tool_call_id: None,
},
finish_reason: FinishReason::ToolCalls,
usage: Usage {
input_tokens: 10,
output_tokens: 5,
total_tokens: 15,
..Default::default()
},
raw: None,
warnings: vec![],
rate_limit: None,
};
let events = vec![
Ok(StreamEvent::ToolCallEnd { tool_call }),
Ok(StreamEvent::finish(
FinishReason::ToolCalls,
response.usage.clone(),
response,
)),
];
Ok(Box::pin(stream::iter(events)))
} else {
// Second stream: return text
let text = "The weather in SF is 72F";
let response = Response {
id: "resp_2".into(),
model: "mock-model".into(),
provider: "mock".into(),
message: Message::assistant(text),
finish_reason: FinishReason::Stop,
usage: Usage {
input_tokens: 20,
output_tokens: 10,
total_tokens: 30,
..Default::default()
},
raw: None,
warnings: vec![],
rate_limit: None,
};
let events = vec![
Ok(StreamEvent::text_delta(text, Some("t1".into()))),
Ok(StreamEvent::finish(
FinishReason::Stop,
response.usage.clone(),
response,
)),
];
Ok(Box::pin(stream::iter(events)))
}
}
}
#[tokio::test]
async fn stream_with_tool_loop_executes_tools() {
let call_count = Arc::new(AtomicU32::new(0));
let provider: Arc<dyn ProviderAdapter> = Arc::new(StreamingToolCallMockProvider {
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 mut result = stream(
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 { Ok(serde_json::json!("72F")) },
)])
.max_tool_rounds(5)
.client(client),
)
.await
.unwrap();
// Collect all events
let mut events = Vec::new();
while let Some(item) = result.next().await {
events.push(item);
}
// Should have events from both rounds
assert_eq!(call_count.load(Ordering::SeqCst), 2);
// Should have text deltas from the second round
let text_deltas: Vec<_> = events
.iter()
.filter_map(|e| match e {
Ok(StreamEvent::TextDelta { delta, .. }) => Some(delta.as_str()),
_ => None,
})
.collect();
assert_eq!(text_deltas, vec!["The weather in SF is 72F"]);
// The final response should be the text response
assert!(result.response().is_some());
assert_eq!(result.response().unwrap().text(), "The weather in SF is 72F");
}
#[tokio::test]
async fn stream_no_tool_loop_when_max_rounds_zero() {
let call_count = Arc::new(AtomicU32::new(0));
let provider: Arc<dyn ProviderAdapter> = Arc::new(StreamingToolCallMockProvider {
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 mut result = stream(
GenerateParams::new("mock-model")
.prompt("What's the weather?")
.tools(vec![Tool::active(
"get_weather",
"Get weather",
serde_json::json!({"type": "object", "properties": {"city": {"type": "string"}}}),
|_args| async { Ok(serde_json::json!("72F")) },
)])
.max_tool_rounds(0)
.client(client),
)
.await
.unwrap();
// Consume all events
while result.next().await.is_some() {}
// Only one stream call, no tool execution
assert_eq!(call_count.load(Ordering::SeqCst), 1);
}
}

View file

@ -8,3 +8,5 @@ pub mod retry;
pub mod generate;
pub mod catalog;
pub mod providers;
pub use tokio_util::sync::CancellationToken;

View file

@ -328,7 +328,10 @@ fn translate_response_format(
}
/// Build the Gemini API request body from a unified `Request`.
fn build_api_request(request: &Request) -> ApiRequest {
///
/// Returns a `serde_json::Value` so that `provider_options.gemini` fields can be
/// merged into the request before sending.
fn build_api_request(request: &Request) -> serde_json::Value {
let (system_text, other_messages) = extract_system_prompt(&request.messages);
let system_instruction = system_text.map(|text| SystemInstruction {
@ -354,12 +357,37 @@ fn build_api_request(request: &Request) -> ApiRequest {
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 {
let api_request = ApiRequest {
contents,
system_instruction,
generation_config: Some(generation_config),
tools: api_tools,
tool_config,
};
let mut body = serde_json::to_value(&api_request).unwrap_or_default();
merge_provider_options(&mut body, request.provider_options.as_ref());
body
}
/// Merge `provider_options.gemini` fields into the serialized API request body.
///
/// Known fields like `safety_settings` and `cached_content` are set directly.
/// Any other fields are merged at the top level, allowing pass-through of
/// Gemini-specific options not covered by the unified schema.
fn merge_provider_options(body: &mut serde_json::Value, provider_options: Option<&serde_json::Value>) {
let Some(gemini_opts) = provider_options.and_then(|opts| opts.get("gemini")) else {
return;
};
let Some(body_map) = body.as_object_mut() else {
return;
};
let Some(gemini_map) = gemini_opts.as_object() else {
return;
};
for (key, value) in gemini_map {
body_map.insert(key.clone(), value.clone());
}
}
@ -688,7 +716,7 @@ impl ProviderAdapter for Adapter {
}
async fn complete(&self, request: &Request) -> Result<Response, SdkError> {
let api_request = build_api_request(request);
let api_body = build_api_request(request);
let url = format!(
"{}/models/{}:generateContent?key={}",
@ -696,7 +724,7 @@ impl ProviderAdapter for Adapter {
);
let body = send_and_read_body(
self.client.post(&url).json(&api_request).timeout(self.request_timeout),
self.client.post(&url).json(&api_body).timeout(self.request_timeout),
"gemini",
"status",
)
@ -751,7 +779,7 @@ impl ProviderAdapter for Adapter {
}
async fn stream(&self, request: &Request) -> Result<StreamEventStream, SdkError> {
let api_request = build_api_request(request);
let api_body = build_api_request(request);
let url = format!(
"{}/models/{}:streamGenerateContent?alt=sse&key={}",
@ -759,8 +787,139 @@ impl ProviderAdapter for Adapter {
);
let http_resp =
send_streaming_request(self.client.post(&url).json(&api_request)).await?;
send_streaming_request(self.client.post(&url).json(&api_body)).await?;
Ok(process_sse_stream(http_resp, request.model.clone()))
}
}
#[cfg(test)]
mod tests {
use super::*;
fn minimal_request() -> Request {
Request {
model: "gemini-2.0-flash".to_string(),
messages: vec![Message::user("Hello")],
provider: None,
tools: None,
tool_choice: None,
response_format: None,
temperature: None,
top_p: None,
max_tokens: None,
stop_sequences: None,
reasoning_effort: None,
metadata: None,
provider_options: None,
}
}
#[test]
fn provider_options_none_produces_standard_body() {
let request = minimal_request();
let body = build_api_request(&request);
assert!(body.get("safetySettings").is_none());
assert!(body.get("cachedContent").is_none());
}
#[test]
fn provider_options_gemini_safety_settings_merged() {
let mut request = minimal_request();
request.provider_options = Some(serde_json::json!({
"gemini": {
"safetySettings": [
{"category": "HARM_CATEGORY_HARASSMENT", "threshold": "BLOCK_NONE"}
]
}
}));
let body = build_api_request(&request);
let safety = body.get("safetySettings").expect("safetySettings should be present");
let arr = safety.as_array().expect("should be an array");
assert_eq!(arr.len(), 1);
assert_eq!(arr[0]["category"], "HARM_CATEGORY_HARASSMENT");
}
#[test]
fn provider_options_gemini_cached_content_merged() {
let mut request = minimal_request();
request.provider_options = Some(serde_json::json!({
"gemini": {
"cachedContent": "projects/my-project/cachedContents/abc123"
}
}));
let body = build_api_request(&request);
assert_eq!(
body.get("cachedContent").and_then(serde_json::Value::as_str),
Some("projects/my-project/cachedContents/abc123")
);
}
#[test]
fn provider_options_gemini_multiple_fields_merged() {
let mut request = minimal_request();
request.provider_options = Some(serde_json::json!({
"gemini": {
"safetySettings": [{"category": "HARM_CATEGORY_HATE_SPEECH", "threshold": "BLOCK_LOW_AND_ABOVE"}],
"cachedContent": "cache-id",
"customField": "custom-value"
}
}));
let body = build_api_request(&request);
assert!(body.get("safetySettings").is_some());
assert_eq!(
body.get("cachedContent").and_then(serde_json::Value::as_str),
Some("cache-id")
);
assert_eq!(
body.get("customField").and_then(serde_json::Value::as_str),
Some("custom-value")
);
}
#[test]
fn provider_options_other_provider_ignored() {
let mut request = minimal_request();
request.provider_options = Some(serde_json::json!({
"anthropic": {
"auto_cache": false
}
}));
let body = build_api_request(&request);
assert!(body.get("auto_cache").is_none());
}
#[test]
fn provider_options_gemini_preserves_standard_fields() {
let mut request = minimal_request();
request.temperature = Some(0.5);
request.max_tokens = Some(100);
request.provider_options = Some(serde_json::json!({
"gemini": {
"cachedContent": "cache-id"
}
}));
let body = build_api_request(&request);
let gen_config = body.get("generationConfig").expect("generationConfig should exist");
assert_eq!(gen_config.get("temperature").and_then(serde_json::Value::as_f64), Some(0.5));
assert_eq!(gen_config.get("maxOutputTokens").and_then(serde_json::Value::as_i64), Some(100));
assert_eq!(
body.get("cachedContent").and_then(serde_json::Value::as_str),
Some("cache-id")
);
}
#[test]
fn merge_provider_options_with_non_object_gemini_value() {
let mut body = serde_json::json!({"contents": []});
let opts = serde_json::json!({"gemini": "not-an-object"});
merge_provider_options(&mut body, Some(&opts));
// Should not crash and body should be unchanged
assert!(body.get("contents").is_some());
}
}

View file

@ -67,6 +67,8 @@ struct ApiRequest {
reasoning: Option<serde_json::Value>,
#[serde(skip_serializing_if = "Option::is_none")]
text: Option<serde_json::Value>,
#[serde(skip_serializing_if = "Option::is_none")]
metadata: Option<std::collections::HashMap<String, String>>,
#[serde(skip_serializing_if = "std::ops::Not::not")]
stream: bool,
}
@ -294,10 +296,31 @@ fn build_api_request(request: &Request, stream: bool) -> ApiRequest {
tool_choice,
reasoning,
text,
metadata: request.metadata.clone(),
stream,
}
}
/// Serialize an `ApiRequest` to JSON and merge any `provider_options.openai` keys into it.
fn build_request_body(request: &Request, stream: bool) -> serde_json::Value {
let api_request = build_api_request(request, stream);
let mut body = serde_json::to_value(&api_request).unwrap_or_else(|_| serde_json::json!({}));
if let Some(openai_opts) = request
.provider_options
.as_ref()
.and_then(|opts| opts.get("openai"))
{
if let (Some(base), Some(overrides)) = (body.as_object_mut(), openai_opts.as_object()) {
for (key, value) in overrides {
base.insert(key.clone(), value.clone());
}
}
}
body
}
/// 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();
@ -744,14 +767,14 @@ impl ProviderAdapter for Adapter {
}
async fn complete(&self, request: &Request) -> Result<Response, SdkError> {
let api_request = build_api_request(request, false);
let request_body = build_request_body(request, false);
let url = format!("{}/responses", self.base_url);
let (body, headers) = send_and_read_response(
self.client
.post(&url)
.bearer_auth(&self.api_key)
.json(&api_request)
.json(&request_body)
.timeout(self.request_timeout),
"openai",
"type",
@ -803,14 +826,14 @@ impl ProviderAdapter for Adapter {
}
async fn stream(&self, request: &Request) -> Result<StreamEventStream, SdkError> {
let api_request = build_api_request(request, true);
let request_body = build_request_body(request, true);
let url = format!("{}/responses", self.base_url);
let http_resp = self
.client
.post(&url)
.bearer_auth(&self.api_key)
.json(&api_request)
.json(&request_body)
.send()
.await
.map_err(|e| SdkError::Network {
@ -868,3 +891,127 @@ impl ProviderAdapter for Adapter {
Ok(Box::pin(stream))
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::collections::HashMap;
fn minimal_request() -> Request {
Request {
model: "gpt-4o".to_string(),
messages: vec![Message::user("Hello")],
provider: None,
tools: None,
tool_choice: None,
response_format: None,
temperature: None,
top_p: None,
max_tokens: None,
stop_sequences: None,
reasoning_effort: None,
metadata: None,
provider_options: None,
}
}
#[test]
fn build_request_body_includes_metadata() {
let mut metadata = HashMap::new();
metadata.insert("user_id".to_string(), "u123".to_string());
metadata.insert("session".to_string(), "s456".to_string());
let mut request = minimal_request();
request.metadata = Some(metadata);
let body = build_request_body(&request, false);
let meta = body.get("metadata").expect("metadata should be present");
assert_eq!(meta["user_id"], "u123");
assert_eq!(meta["session"], "s456");
}
#[test]
fn build_request_body_omits_metadata_when_none() {
let request = minimal_request();
let body = build_request_body(&request, false);
assert!(body.get("metadata").is_none());
}
#[test]
fn build_request_body_merges_provider_options_openai() {
let mut request = minimal_request();
request.provider_options = Some(serde_json::json!({
"openai": {
"store": true,
"previous_response_id": "resp_abc123"
}
}));
let body = build_request_body(&request, false);
assert_eq!(body["store"], true);
assert_eq!(body["previous_response_id"], "resp_abc123");
}
#[test]
fn build_request_body_provider_options_override_fields() {
let mut request = minimal_request();
request.temperature = Some(0.5);
request.provider_options = Some(serde_json::json!({
"openai": {
"temperature": 0.9
}
}));
let body = build_request_body(&request, false);
// provider_options should override the base field
assert_eq!(body["temperature"], 0.9);
}
#[test]
fn build_request_body_ignores_non_openai_provider_options() {
let mut request = minimal_request();
request.provider_options = Some(serde_json::json!({
"anthropic": {
"thinking": {"type": "enabled", "budget_tokens": 10000}
}
}));
let body = build_request_body(&request, false);
// anthropic options should not leak into the OpenAI request
assert!(body.get("thinking").is_none());
}
#[test]
fn build_request_body_no_provider_options() {
let request = minimal_request();
let body = build_request_body(&request, false);
assert_eq!(body["model"], "gpt-4o");
// stream field is omitted when false (skip_serializing_if)
assert!(body.get("stream").is_none());
}
#[test]
fn build_request_body_stream_flag() {
let request = minimal_request();
let body = build_request_body(&request, true);
assert!(body["stream"].as_bool().unwrap_or(false));
}
#[test]
fn build_request_body_metadata_and_provider_options_together() {
let mut metadata = HashMap::new();
metadata.insert("trace_id".to_string(), "t789".to_string());
let mut request = minimal_request();
request.metadata = Some(metadata);
request.provider_options = Some(serde_json::json!({
"openai": {
"store": true
}
}));
let body = build_request_body(&request, false);
assert_eq!(body["metadata"]["trace_id"], "t789");
assert_eq!(body["store"], true);
}
}

View file

@ -291,8 +291,11 @@ fn translate_response_format(format: &ResponseFormat) -> serde_json::Value {
}
}
/// Build an `ApiRequest` from a unified `Request`.
fn build_api_request(request: &Request, stream: Option<bool>) -> ApiRequest {
/// Build the API request body from a unified `Request`.
///
/// Returns a `serde_json::Value` so that `provider_options.<provider_name>` fields
/// can be merged into the request before sending.
fn build_api_request(request: &Request, stream: Option<bool>, provider_name: &str) -> serde_json::Value {
let chat_messages = translate_messages(&request.messages);
let tools = request.tools.as_ref().map(|t| translate_tools(t));
let tool_choice = request.tool_choice.as_ref().map(translate_tool_choice);
@ -301,7 +304,7 @@ fn build_api_request(request: &Request, stream: Option<bool>) -> ApiRequest {
.as_ref()
.map(translate_response_format);
ApiRequest {
let api_request = ApiRequest {
model: request.model.clone(),
messages: chat_messages,
temperature: request.temperature,
@ -312,6 +315,34 @@ fn build_api_request(request: &Request, stream: Option<bool>) -> ApiRequest {
tool_choice,
response_format,
stream,
};
let mut body = serde_json::to_value(&api_request).unwrap_or_default();
merge_provider_options(&mut body, request.provider_options.as_ref(), provider_name);
body
}
/// Merge `provider_options.<provider_name>` fields into the serialized API request body.
///
/// The provider name is configurable (e.g. "groq", "together", "openai-compatible"),
/// allowing each instance to have its own namespace in `provider_options`.
fn merge_provider_options(
body: &mut serde_json::Value,
provider_options: Option<&serde_json::Value>,
provider_name: &str,
) {
let Some(opts) = provider_options.and_then(|opts| opts.get(provider_name)) else {
return;
};
let Some(body_map) = body.as_object_mut() else {
return;
};
let Some(opts_map) = opts.as_object() else {
return;
};
for (key, value) in opts_map {
body_map.insert(key.clone(), value.clone());
}
}
@ -323,14 +354,14 @@ impl ProviderAdapter for Adapter {
}
async fn complete(&self, request: &Request) -> Result<Response, SdkError> {
let api_request = build_api_request(request, None);
let api_body = build_api_request(request, None, &self.provider_name);
let url = format!("{}/chat/completions", self.base_url);
let (body, headers) = send_and_read_response(
self.client
.post(&url)
.bearer_auth(&self.api_key)
.json(&api_request)
.json(&api_body)
.timeout(self.request_timeout),
&self.provider_name,
"type",
@ -397,14 +428,14 @@ impl ProviderAdapter for Adapter {
}
async fn stream(&self, request: &Request) -> Result<StreamEventStream, SdkError> {
let api_request = build_api_request(request, Some(true));
let api_body = build_api_request(request, Some(true), &self.provider_name);
let url = format!("{}/chat/completions", self.base_url);
let http_resp = self
.client
.post(&url)
.bearer_auth(&self.api_key)
.json(&api_request)
.json(&api_body)
.send()
.await
.map_err(|e| SdkError::Network {
@ -1111,4 +1142,111 @@ mod tests {
assert_eq!(tool_calls[0]["id"], "call_1");
assert_eq!(tool_calls[0]["function"]["name"], "get_weather");
}
fn minimal_request() -> Request {
Request {
model: "llama-3.1-70b".to_string(),
messages: vec![Message::user("Hello")],
provider: None,
tools: None,
tool_choice: None,
response_format: None,
temperature: None,
top_p: None,
max_tokens: None,
stop_sequences: None,
reasoning_effort: None,
metadata: None,
provider_options: None,
}
}
#[test]
fn provider_options_none_produces_standard_body() {
let request = minimal_request();
let body = build_api_request(&request, None, "groq");
assert_eq!(body["model"], "llama-3.1-70b");
assert!(body.get("stream").is_none());
}
#[test]
fn provider_options_matching_name_merged() {
let mut request = minimal_request();
request.provider_options = Some(serde_json::json!({
"groq": {
"frequency_penalty": 0.5,
"presence_penalty": 0.3
}
}));
let body = build_api_request(&request, None, "groq");
assert_eq!(body["frequency_penalty"], 0.5);
assert_eq!(body["presence_penalty"], 0.3);
}
#[test]
fn provider_options_different_name_ignored() {
let mut request = minimal_request();
request.provider_options = Some(serde_json::json!({
"together": {
"repetition_penalty": 1.2
}
}));
let body = build_api_request(&request, None, "groq");
assert!(body.get("repetition_penalty").is_none());
}
#[test]
fn provider_options_uses_adapter_name() {
let mut request = minimal_request();
request.provider_options = Some(serde_json::json!({
"together": {
"repetition_penalty": 1.2
}
}));
let body = build_api_request(&request, None, "together");
assert_eq!(body["repetition_penalty"], 1.2);
}
#[test]
fn provider_options_preserves_standard_fields() {
let mut request = minimal_request();
request.temperature = Some(0.7);
request.max_tokens = Some(200);
request.provider_options = Some(serde_json::json!({
"groq": {
"frequency_penalty": 0.5
}
}));
let body = build_api_request(&request, Some(true), "groq");
assert_eq!(body["temperature"], 0.7);
assert_eq!(body["max_tokens"], 200);
assert_eq!(body["stream"], true);
assert_eq!(body["frequency_penalty"], 0.5);
}
#[test]
fn provider_options_can_override_model() {
let mut request = minimal_request();
request.provider_options = Some(serde_json::json!({
"groq": {
"model": "custom-model"
}
}));
let body = build_api_request(&request, None, "groq");
assert_eq!(body["model"], "custom-model");
}
#[test]
fn merge_provider_options_with_non_object_value() {
let mut body = serde_json::json!({"model": "test"});
let opts = serde_json::json!({"groq": "not-an-object"});
merge_provider_options(&mut body, Some(&opts), "groq");
// Should not crash and body should be unchanged
assert_eq!(body["model"], "test");
}
}