mirror of
https://github.com/fabro-sh/fabro.git
synced 2026-08-28 05:27:41 +00:00
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:
parent
ed07d43335
commit
5f2385b87b
8 changed files with 1134 additions and 23 deletions
1
Cargo.lock
generated
1
Cargo.lock
generated
|
|
@ -1401,6 +1401,7 @@ dependencies = [
|
|||
"thiserror",
|
||||
"tokio",
|
||||
"tokio-stream",
|
||||
"tokio-util",
|
||||
"uuid",
|
||||
]
|
||||
|
||||
|
|
|
|||
|
|
@ -29,3 +29,4 @@ tokio-stream = "0.1"
|
|||
async-trait = "0.1"
|
||||
base64 = "0.22"
|
||||
bytes = "1"
|
||||
tokio-util = "0.7"
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(¶ms, &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(¶ms)?;
|
||||
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, ¶ms, &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(¶ms, &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(¶ms, &messages, tool_definitions.as_deref());
|
||||
client.stream(&request).await
|
||||
stream_generate_raw(&client, ¶ms, &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);
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -8,3 +8,5 @@ pub mod retry;
|
|||
pub mod generate;
|
||||
pub mod catalog;
|
||||
pub mod providers;
|
||||
|
||||
pub use tokio_util::sync::CancellationToken;
|
||||
|
|
|
|||
|
|
@ -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());
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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);
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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");
|
||||
}
|
||||
}
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue