Simplify LLM crate: constructors, shared HTTP/SSE infra, Deref delegation

- Add ToolResult::success()/error() constructors, replacing 16 manual
  construction sites across tools.rs, session.rs, history.rs, generate.rs
- Collapse ContentPart::RedactedThinking into Thinking (use redacted field)
- Extract HttpApi base struct shared by all 4 provider adapters
- Extract LineReader into common.rs for shared SSE byte buffering and
  timeout handling; convert Anthropic/OpenAI from BoxStream to Response
- Replace manual delegation methods on GenerateResult/StepResult with
  Deref<Target=Response>

Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
This commit is contained in:
Bryan Helmkamp 2026-02-27 19:12:33 -05:00
parent 0756a42ce8
commit 3253589193
13 changed files with 304 additions and 571 deletions

View file

@ -51,7 +51,7 @@ impl History {
// doesn't already contain thinking blocks (which preserve signatures).
let has_thinking_parts = provider_parts
.iter()
.any(|p| matches!(p, ContentPart::Thinking(_) | ContentPart::RedactedThinking(_)));
.any(|p| matches!(p, ContentPart::Thinking(_)));
if !has_thinking_parts {
if let Some(reasoning_text) = reasoning {
parts.push(ContentPart::Thinking(
@ -308,13 +308,7 @@ mod tests {
#[test]
fn tool_results_turn_maps_to_tool_message() {
let mut history = History::default();
let result = ToolResult {
tool_call_id: "call_1".into(),
content: serde_json::json!("file contents here"),
is_error: false,
image_data: None,
image_media_type: None,
};
let result = ToolResult::success("call_1", serde_json::json!("file contents here"));
history.push(Turn::ToolResults {
results: vec![result],
timestamp: SystemTime::now(),
@ -398,13 +392,7 @@ mod tests {
timestamp: SystemTime::now(),
});
history.push(Turn::ToolResults {
results: vec![ToolResult {
tool_call_id: "c1".into(),
content: serde_json::json!("file1.rs\nfile2.rs"),
is_error: false,
image_data: None,
image_media_type: None,
}],
results: vec![ToolResult::success("c1", serde_json::json!("file1.rs\nfile2.rs"))],
timestamp: SystemTime::now(),
});

View file

@ -408,7 +408,6 @@ impl Session {
p,
llm::types::ContentPart::Other { .. }
| llm::types::ContentPart::Thinking(_)
| llm::types::ContentPart::RedactedThinking(_)
)
})
.cloned()
@ -632,13 +631,7 @@ and conversational filler.".to_string()),
let mut results = Vec::new();
for tc in tool_calls {
if self.cancel_token.is_cancelled() {
results.push(ToolResult {
tool_call_id: tc.id.clone(),
content: serde_json::json!("Cancelled"),
is_error: true,
image_data: None,
image_media_type: None,
});
results.push(ToolResult::error(tc.id.clone(), "Cancelled"));
continue;
}
@ -821,13 +814,7 @@ async fn execute_one_tool(
) -> ToolResult {
if let Some(approval_fn) = tool_approval {
if let Err(denial_message) = approval_fn(tool_name, arguments) {
return ToolResult {
tool_call_id: tool_call_id.to_string(),
content: serde_json::json!(denial_message),
is_error: true,
image_data: None,
image_media_type: None,
};
return ToolResult::error(tool_call_id, denial_message);
}
}
@ -836,39 +823,15 @@ async fn execute_one_tool(
if let Err(validation_error) =
validate_tool_args(&registered_tool.definition.parameters, arguments)
{
return ToolResult {
tool_call_id: tool_call_id.to_string(),
content: serde_json::json!(validation_error),
is_error: true,
image_data: None,
image_media_type: None,
};
return ToolResult::error(tool_call_id, validation_error);
}
match (registered_tool.executor)(arguments.clone(), env, cancel_token).await {
Ok(output) => ToolResult {
tool_call_id: tool_call_id.to_string(),
content: serde_json::json!(output),
is_error: false,
image_data: None,
image_media_type: None,
},
Err(err) => ToolResult {
tool_call_id: tool_call_id.to_string(),
content: serde_json::json!(err),
is_error: true,
image_data: None,
image_media_type: None,
},
Ok(output) => ToolResult::success(tool_call_id, serde_json::json!(output)),
Err(err) => ToolResult::error(tool_call_id, err),
}
}
None => ToolResult {
tool_call_id: tool_call_id.to_string(),
content: serde_json::json!(format!("Unknown tool: {tool_name}")),
is_error: true,
image_data: None,
image_media_type: None,
},
None => ToolResult::error(tool_call_id, format!("Unknown tool: {tool_name}")),
}
}

View file

@ -222,14 +222,14 @@ async fn run_prompt(args: PromptArgs) -> Result<()> {
let object = result.output.as_ref().unwrap_or(&serde_json::Value::Null);
println!("{}", serde_json::to_string_pretty(object)?);
if args.usage {
print_usage(result.usage());
print_usage(&result.usage);
}
}
(true, None) => {
let result = generate::generate(params).await?;
print!("{}", result.text());
if args.usage {
print_usage(result.usage());
print_usage(&result.usage);
}
}
(false, Some(schema)) => {

View file

@ -1145,8 +1145,8 @@ mod tests {
.unwrap();
assert_eq!(result.text(), "Hi there!");
assert_eq!(*result.finish_reason(), FinishReason::Stop);
assert_eq!(result.usage().input_tokens, 10);
assert_eq!(result.finish_reason, FinishReason::Stop);
assert_eq!(result.usage.input_tokens, 10);
assert_eq!(result.steps.len(), 1);
}
@ -2164,13 +2164,7 @@ mod tests {
serde_json::json!({"city": "SF"}),
)];
let tool_results = vec![crate::types::ToolResult {
tool_call_id: "call_1".into(),
content: serde_json::json!("72F"),
is_error: false,
image_data: None,
image_media_type: None,
}];
let tool_results = vec![crate::types::ToolResult::success("call_1", serde_json::json!("72F"))];
// Processing StepFinish should not panic and should not set the final response
acc.process(&StreamEvent::step_finish(

View file

@ -13,57 +13,35 @@ use crate::types::{
/// Provider adapter for the Anthropic Messages API.
pub struct Adapter {
api_key: String,
base_url: String,
default_headers: std::collections::HashMap<String, String>,
client: reqwest::Client,
request_timeout: Option<std::time::Duration>,
stream_read_timeout: Option<std::time::Duration>,
pub(crate) http: super::http_api::HttpApi,
}
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(),
base_url: DEFAULT_BASE_URL.to_string(),
default_headers: std::collections::HashMap::new(),
client,
request_timeout: timeout.request.map(std::time::Duration::from_secs_f64),
stream_read_timeout: timeout.stream_read.map(std::time::Duration::from_secs_f64),
http: super::http_api::HttpApi::new(api_key, DEFAULT_BASE_URL),
}
}
#[must_use]
pub fn with_base_url(mut self, base_url: impl Into<String>) -> Self {
self.base_url = base_url.into();
self.http.base_url = base_url.into();
self
}
#[must_use]
pub fn with_default_headers(mut self, headers: std::collections::HashMap<String, String>) -> Self {
self.default_headers = headers;
self
pub fn with_default_headers(self, headers: std::collections::HashMap<String, String>) -> Self {
Self { http: self.http.with_default_headers(headers) }
}
#[must_use]
pub fn with_timeout(mut self, timeout: crate::types::AdapterTimeout) -> Self {
self.client = reqwest::Client::builder()
.connect_timeout(std::time::Duration::from_secs_f64(timeout.connect))
.build()
.unwrap_or_default();
self.request_timeout = timeout.request.map(std::time::Duration::from_secs_f64);
self.stream_read_timeout = timeout.stream_read.map(std::time::Duration::from_secs_f64);
self
pub fn with_timeout(self, timeout: crate::types::AdapterTimeout) -> Self {
Self { http: self.http.with_timeout(timeout) }
}
fn messages_url(&self) -> String {
format!("{}/messages", self.base_url)
format!("{}/messages", self.http.base_url)
}
}
@ -196,7 +174,7 @@ fn parse_content_block(block: &serde_json::Value) -> Option<ContentPart> {
.map(String::from),
redacted: false,
})),
"redacted_thinking" => Some(ContentPart::RedactedThinking(ThinkingData {
"redacted_thinking" => Some(ContentPart::Thinking(ThinkingData {
text: block
.get("data")
.and_then(serde_json::Value::as_str)
@ -231,6 +209,10 @@ fn content_part_to_api(part: &ContentPart) -> Option<serde_json::Value> {
"is_error": tr.is_error,
}))
}
ContentPart::Thinking(td) if td.redacted => Some(serde_json::json!({
"type": "redacted_thinking",
"data": td.text,
})),
ContentPart::Thinking(td) => {
let mut block = serde_json::json!({
"type": "thinking",
@ -241,10 +223,6 @@ fn content_part_to_api(part: &ContentPart) -> Option<serde_json::Value> {
}
Some(block)
}
ContentPart::RedactedThinking(td) => Some(serde_json::json!({
"type": "redacted_thinking",
"data": td.text,
})),
ContentPart::Image(img) => {
if let Some(url) = &img.url {
if crate::providers::common::is_file_path(url) {
@ -918,34 +896,25 @@ enum SseResult {
}
struct SseReaderState {
byte_stream: futures::stream::BoxStream<'static, Result<bytes::Bytes, reqwest::Error>>,
buffer: String,
line_reader: super::common::LineReader,
accumulator: StreamAccumulator,
pending_events: std::collections::VecDeque<StreamEvent>,
done: bool,
/// When true, `tool_use` events for the synthetic tool are converted to text events.
json_schema_mode: bool,
stream_read_timeout: Option<std::time::Duration>,
}
impl SseReaderState {
fn new(
byte_stream: impl futures::Stream<Item = Result<bytes::Bytes, reqwest::Error>>
+ Send
+ 'static,
http_resp: reqwest::Response,
rate_limit: Option<crate::types::RateLimitInfo>,
json_schema_mode: bool,
stream_read_timeout: Option<std::time::Duration>,
) -> Self {
use futures::StreamExt;
Self {
byte_stream: byte_stream.boxed(),
buffer: String::new(),
line_reader: super::common::LineReader::new(http_resp, stream_read_timeout),
accumulator: StreamAccumulator::new(rate_limit),
pending_events: std::collections::VecDeque::new(),
done: false,
json_schema_mode,
stream_read_timeout,
}
}
@ -954,59 +923,24 @@ impl SseReaderState {
/// SSE events are separated by double newlines. Each event has optional
/// `event:` and `data:` lines.
async fn next_sse_event(&mut self) -> SseResult {
use futures::StreamExt;
loop {
// Try to extract a complete SSE event from the buffer.
if let Some(result) = self.try_parse_event() {
return result;
}
if self.done {
return SseResult::Done;
}
// Read more bytes from the stream.
let chunk_result = match self.stream_read_timeout {
Some(timeout) => tokio::time::timeout(timeout, self.byte_stream.next()).await,
None => Ok(self.byte_stream.next().await),
};
match chunk_result {
Ok(Some(Ok(chunk))) => {
let text = String::from_utf8_lossy(&chunk);
self.buffer.push_str(&text);
}
Ok(Some(Err(e))) => {
return SseResult::Error(SdkError::Stream {
message: e.to_string(),
});
}
Ok(None) => {
self.done = true;
// Try one more time to parse any remaining data.
if let Some(result) = self.try_parse_event() {
match self.line_reader.read_next_chunk("\n\n").await {
Ok(Some(event_block)) => {
if let Some(result) = Self::parse_event_block(&event_block) {
return result;
}
return SseResult::Done;
}
Err(_) => {
return SseResult::Error(SdkError::Stream {
message: "stream read timed out waiting for next event".to_string(),
});
// No data in this block (e.g. heartbeat comment); keep reading.
}
Ok(None) => return SseResult::Done,
Err(e) => return SseResult::Error(e),
}
}
}
/// Attempt to parse one complete SSE event from the buffer.
/// Parse an SSE event block into an `SseResult`.
///
/// Returns `None` if no complete event is available yet.
fn try_parse_event(&mut self) -> Option<SseResult> {
// SSE events are terminated by a blank line (double newline).
let separator = self.buffer.find("\n\n")?;
let event_block = self.buffer[..separator].to_string();
self.buffer = self.buffer[separator + 2..].to_string();
/// Returns `None` for blocks with no `data:` lines (e.g. heartbeat comments).
fn parse_event_block(event_block: &str) -> Option<SseResult> {
let mut event_type = String::new();
let mut data_parts: Vec<String> = Vec::new();
@ -1117,13 +1051,13 @@ fn build_api_request(
};
let url = adapter.messages_url();
let mut req_builder = adapter.client.post(&url);
let mut req_builder = adapter.http.client.post(&url);
// Apply default_headers first so adapter-specific headers can override
for (key, value) in &adapter.default_headers {
for (key, value) in &adapter.http.default_headers {
req_builder = req_builder.header(key, value);
}
req_builder = req_builder
.header("x-api-key", &adapter.api_key)
.header("x-api-key", &adapter.http.api_key)
.header("anthropic-version", "2023-06-01");
if let Some(beta_str) = build_beta_header(request.provider_options.as_ref(), auto_cache) {
@ -1147,7 +1081,7 @@ impl ProviderAdapter for Adapter {
let (_api_request, req_builder) = build_api_request(self, request, false);
let mut req = req_builder;
if let Some(t) = self.request_timeout {
if let Some(t) = self.http.request_timeout {
req = req.timeout(t);
}
let (body, headers) =
@ -1235,12 +1169,11 @@ impl ProviderAdapter for Adapter {
}
let rate_limit = parse_rate_limit_headers(http_resp.headers());
let byte_stream = http_resp.bytes_stream();
let json_schema_mode = uses_json_schema_format(request);
let stream_read_timeout = self.stream_read_timeout;
let stream_read_timeout = self.http.stream_read_timeout;
let stream = futures::stream::unfold(
SseReaderState::new(byte_stream, rate_limit, json_schema_mode, stream_read_timeout),
SseReaderState::new(http_resp, rate_limit, json_schema_mode, stream_read_timeout),
|mut state| async move {
loop {
// Drain any buffered events first.

View file

@ -198,6 +198,71 @@ pub async fn send_and_read_response(
Ok((body, headers))
}
/// Shared line reader for SSE streams.
///
/// Buffers bytes from a `reqwest::Response` and splits them by a configurable
/// delimiter (e.g. `"\n"` for Gemini/OpenAI-compatible, `"\n\n"` for
/// Anthropic/OpenAI SSE event blocks).
pub struct LineReader {
response: reqwest::Response,
buffer: String,
stream_read_timeout: Option<std::time::Duration>,
}
impl LineReader {
pub fn new(response: reqwest::Response, stream_read_timeout: Option<std::time::Duration>) -> Self {
Self {
response,
buffer: String::new(),
stream_read_timeout,
}
}
/// Read the next complete segment delimited by `delimiter`.
///
/// Returns `Ok(Some(segment))` for each complete segment, `Ok(None)` when
/// the stream is exhausted, or `Err` on I/O or timeout errors. When the
/// stream ends with data remaining in the buffer, the leftover is returned
/// as a final segment.
pub async fn read_next_chunk(&mut self, delimiter: &str) -> Result<Option<String>, SdkError> {
loop {
if let Some(pos) = self.buffer.find(delimiter) {
let segment = self.buffer[..pos].to_string();
self.buffer = self.buffer[pos + delimiter.len()..].to_string();
return Ok(Some(segment));
}
let chunk_result = match self.stream_read_timeout {
Some(timeout) => tokio::time::timeout(timeout, self.response.chunk()).await,
None => Ok(self.response.chunk().await),
};
match chunk_result {
Ok(Ok(Some(bytes))) => {
let text = String::from_utf8_lossy(&bytes);
self.buffer.push_str(&text);
}
Ok(Ok(None)) => {
if self.buffer.is_empty() {
return Ok(None);
}
let remaining = std::mem::take(&mut self.buffer);
return Ok(Some(remaining));
}
Ok(Err(e)) => {
return Err(SdkError::Stream {
message: e.to_string(),
});
}
Err(_) => {
return Err(SdkError::Stream {
message: "stream read timed out waiting for next event".to_string(),
});
}
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;

View file

@ -15,53 +15,31 @@ 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,
default_headers: std::collections::HashMap<String, String>,
client: reqwest::Client,
request_timeout: Option<std::time::Duration>,
stream_read_timeout: Option<std::time::Duration>,
pub(crate) http: super::http_api::HttpApi,
}
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(),
base_url: DEFAULT_BASE_URL.to_string(),
default_headers: std::collections::HashMap::new(),
client,
request_timeout: timeout.request.map(std::time::Duration::from_secs_f64),
stream_read_timeout: timeout.stream_read.map(std::time::Duration::from_secs_f64),
http: super::http_api::HttpApi::new(api_key, DEFAULT_BASE_URL),
}
}
#[must_use]
pub fn with_base_url(mut self, base_url: impl Into<String>) -> Self {
self.base_url = base_url.into();
self.http.base_url = base_url.into();
self
}
#[must_use]
pub fn with_default_headers(mut self, headers: std::collections::HashMap<String, String>) -> Self {
self.default_headers = headers;
self
pub fn with_default_headers(self, headers: std::collections::HashMap<String, String>) -> Self {
Self { http: self.http.with_default_headers(headers) }
}
#[must_use]
pub fn with_timeout(mut self, timeout: crate::types::AdapterTimeout) -> Self {
self.client = reqwest::Client::builder()
.connect_timeout(std::time::Duration::from_secs_f64(timeout.connect))
.build()
.unwrap_or_default();
self.request_timeout = timeout.request.map(std::time::Duration::from_secs_f64);
self.stream_read_timeout = timeout.stream_read.map(std::time::Duration::from_secs_f64);
self
pub fn with_timeout(self, timeout: crate::types::AdapterTimeout) -> Self {
Self { http: self.http.with_timeout(timeout) }
}
}
@ -652,10 +630,8 @@ fn process_sse_stream(http_resp: reqwest::Response, model: String, rate_limit: O
/// Internal state for the SSE stream processor.
struct SseStreamState {
http_resp: reqwest::Response,
line_reader: super::common::LineReader,
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.
@ -680,15 +656,13 @@ struct SseStreamState {
finished: bool,
/// Rate limit info parsed from HTTP response headers.
rate_limit: Option<crate::types::RateLimitInfo>,
stream_read_timeout: Option<std::time::Duration>,
}
impl SseStreamState {
fn new(http_resp: reqwest::Response, model: String, rate_limit: Option<crate::types::RateLimitInfo>, stream_read_timeout: Option<std::time::Duration>) -> Self {
Self {
http_resp,
line_reader: super::common::LineReader::new(http_resp, stream_read_timeout),
model,
line_buffer: String::new(),
pending_events: std::collections::VecDeque::new(),
stream_started: false,
text_started: false,
@ -701,7 +675,6 @@ impl SseStreamState {
finish_reason_str: None,
finished: false,
rate_limit,
stream_read_timeout,
}
}
@ -709,50 +682,10 @@ impl SseStreamState {
///
/// 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.
let chunk_result = match self.stream_read_timeout {
Some(timeout) => tokio::time::timeout(timeout, self.http_resp.chunk()).await,
None => Ok(self.http_resp.chunk().await),
};
match chunk_result {
Ok(Ok(Some(bytes))) => {
let text = String::from_utf8_lossy(&bytes);
self.line_buffer.push_str(&text);
}
Ok(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));
}
Ok(Err(e)) => {
return Err(SdkError::Stream {
message: format!("error reading Gemini stream: {e}"),
});
}
Err(_) => {
return Err(SdkError::Stream {
message: "stream read timed out waiting for next event".to_string(),
});
}
}
}
self.line_reader
.read_next_chunk("\n")
.await
.map(|opt| opt.map(|s| s.trim_end_matches('\r').to_string()))
}
/// Extract stream events from a parsed SSE chunk and buffer them.
@ -915,15 +848,15 @@ impl ProviderAdapter for Adapter {
let url = format!(
"{}/models/{}:generateContent?key={}",
self.base_url, request.model, self.api_key
self.http.base_url, request.model, self.http.api_key
);
let mut req = self.client.post(&url);
for (key, value) in &self.default_headers {
let mut req = self.http.client.post(&url);
for (key, value) in &self.http.default_headers {
req = req.header(key, value);
}
let mut gemini_req = req.json(&api_body);
if let Some(t) = self.request_timeout {
if let Some(t) = self.http.request_timeout {
gemini_req = gemini_req.timeout(t);
}
let (body, headers) = send_gemini_response(gemini_req).await?;
@ -984,17 +917,17 @@ impl ProviderAdapter for Adapter {
let url = format!(
"{}/models/{}:streamGenerateContent?alt=sse&key={}",
self.base_url, request.model, self.api_key
self.http.base_url, request.model, self.http.api_key
);
let mut req = self.client.post(&url);
for (key, value) in &self.default_headers {
let mut req = self.http.client.post(&url);
for (key, value) in &self.http.default_headers {
req = req.header(key, value);
}
let http_resp = send_streaming_request(req.json(&api_body)).await?;
let rate_limit = parse_rate_limit_headers(http_resp.headers());
Ok(process_sse_stream(http_resp, request.model.clone(), rate_limit, self.stream_read_timeout))
Ok(process_sse_stream(http_resp, request.model.clone(), rate_limit, self.http.stream_read_timeout))
}
}

View file

@ -0,0 +1,54 @@
use std::collections::HashMap;
use std::time::Duration;
use crate::types::AdapterTimeout;
/// Shared HTTP infrastructure for provider adapters.
///
/// Holds the API key, base URL, reqwest client, default headers, and timeout
/// configuration that every provider needs. Provider-specific fields live on
/// the adapter struct itself.
pub struct HttpApi {
pub(crate) api_key: String,
pub(crate) base_url: String,
pub(crate) default_headers: HashMap<String, String>,
pub(crate) client: reqwest::Client,
pub(crate) request_timeout: Option<Duration>,
pub(crate) stream_read_timeout: Option<Duration>,
}
impl HttpApi {
#[must_use]
pub fn new(api_key: impl Into<String>, base_url: impl Into<String>) -> Self {
let timeout = AdapterTimeout::default();
let client = reqwest::Client::builder()
.connect_timeout(Duration::from_secs_f64(timeout.connect))
.build()
.unwrap_or_default();
Self {
api_key: api_key.into(),
base_url: base_url.into(),
default_headers: HashMap::new(),
client,
request_timeout: timeout.request.map(Duration::from_secs_f64),
stream_read_timeout: timeout.stream_read.map(Duration::from_secs_f64),
}
}
#[must_use]
pub fn with_timeout(mut self, timeout: AdapterTimeout) -> Self {
self.client = reqwest::Client::builder()
.connect_timeout(Duration::from_secs_f64(timeout.connect))
.build()
.unwrap_or_default();
self.request_timeout = timeout.request.map(Duration::from_secs_f64);
self.stream_read_timeout = timeout.stream_read.map(Duration::from_secs_f64);
self
}
#[must_use]
pub fn with_default_headers(mut self, headers: HashMap<String, String>) -> Self {
self.default_headers = headers;
self
}
}

View file

@ -1,6 +1,7 @@
pub mod anthropic;
pub mod common;
pub mod gemini;
pub mod http_api;
pub mod openai;
pub mod openai_compatible;

View file

@ -16,39 +16,24 @@ use crate::types::{
/// 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,
pub(crate) http: super::http_api::HttpApi,
org_id: Option<String>,
project_id: Option<String>,
default_headers: std::collections::HashMap<String, String>,
client: reqwest::Client,
request_timeout: Option<std::time::Duration>,
stream_read_timeout: Option<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(),
base_url: "https://api.openai.com/v1".to_string(),
http: super::http_api::HttpApi::new(api_key, "https://api.openai.com/v1"),
org_id: None,
project_id: None,
default_headers: std::collections::HashMap::new(),
client,
request_timeout: timeout.request.map(std::time::Duration::from_secs_f64),
stream_read_timeout: timeout.stream_read.map(std::time::Duration::from_secs_f64),
}
}
#[must_use]
pub fn with_base_url(mut self, base_url: impl Into<String>) -> Self {
self.base_url = base_url.into();
self.http.base_url = base_url.into();
self
}
@ -65,30 +50,23 @@ impl Adapter {
}
#[must_use]
pub fn with_default_headers(mut self, headers: std::collections::HashMap<String, String>) -> Self {
self.default_headers = headers;
self
pub fn with_default_headers(self, headers: std::collections::HashMap<String, String>) -> Self {
Self { http: self.http.with_default_headers(headers), ..self }
}
#[must_use]
pub fn with_timeout(mut self, timeout: crate::types::AdapterTimeout) -> Self {
self.client = reqwest::Client::builder()
.connect_timeout(std::time::Duration::from_secs_f64(timeout.connect))
.build()
.unwrap_or_default();
self.request_timeout = timeout.request.map(std::time::Duration::from_secs_f64);
self.stream_read_timeout = timeout.stream_read.map(std::time::Duration::from_secs_f64);
self
pub fn with_timeout(self, timeout: crate::types::AdapterTimeout) -> Self {
Self { http: self.http.with_timeout(timeout), ..self }
}
/// Build a `reqwest::RequestBuilder` with default headers, org/project headers, and auth.
fn build_request(&self, url: &str) -> reqwest::RequestBuilder {
let mut req = self.client.post(url);
let mut req = self.http.client.post(url);
// Apply default_headers first so adapter-specific headers can override
for (key, value) in &self.default_headers {
for (key, value) in &self.http.default_headers {
req = req.header(key, value);
}
req = req.bearer_auth(&self.api_key);
req = req.bearer_auth(&self.http.api_key);
if let Some(org_id) = &self.org_id {
req = req.header("OpenAI-Organization", org_id);
}
@ -479,10 +457,7 @@ fn parse_output(output: &[serde_json::Value]) -> (Vec<ContentPart>, bool) {
/// 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,
line_reader: super::common::LineReader,
model: String,
response_id: String,
response_model: String,
@ -496,59 +471,39 @@ struct SseStreamState {
emitted_text_start: bool,
raw_response: Option<serde_json::Value>,
rate_limit: Option<crate::types::RateLimitInfo>,
stream_read_timeout: Option<std::time::Duration>,
}
/// Extract complete SSE messages from the buffer.
/// Parse a single SSE message block into an (`event_type`, `data`) pair.
///
/// 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();
/// Each SSE message consists of one or more lines (`event:` and `data:` prefixed).
/// Returns `None` if the block has no `data:` lines.
fn parse_sse_message(message_block: &str) -> Option<(Option<String>, String)> {
let mut current_event: Option<String> = None;
let mut current_data = String::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());
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');
}
}
if !current_data.is_empty() {
messages.push((current_event, current_data));
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());
}
}
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));
if current_data.is_empty() {
None
} else {
Some((current_event, current_data))
}
events
}
/// Process the next chunk(s) from the byte stream and return `StreamEvent`s.
@ -556,44 +511,17 @@ 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() {
let events = dispatch_sse_messages(state, messages);
if !events.is_empty() {
return Ok(events);
}
// All SSE messages were unhandled event types; continue reading.
continue;
}
let chunk_result = match state.stream_read_timeout {
Some(timeout) => tokio::time::timeout(timeout, state.byte_stream.next()).await,
None => Ok(state.byte_stream.next().await),
};
match chunk_result {
Ok(Some(Ok(bytes))) => {
let text = String::from_utf8_lossy(&bytes);
state.buffer.push_str(&text);
}
Ok(Some(Err(e))) => {
return Err(SdkError::Stream {
message: e.to_string(),
});
}
Ok(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));
match state.line_reader.read_next_chunk("\n\n").await? {
Some(message_block) => {
if let Some((event_type, data)) = parse_sse_message(&message_block) {
let events = process_sse_event(state, event_type.as_deref(), &data);
if !events.is_empty() {
return Ok(events);
}
}
return Ok(vec![]);
}
Err(_) => {
return Err(SdkError::Stream {
message: "stream read timed out waiting for next event".to_string(),
});
// No data or unhandled event type; keep reading.
}
None => return Ok(vec![]),
}
}
}
@ -915,10 +843,10 @@ impl ProviderAdapter for Adapter {
crate::provider::validate_tool_choice(self, tc)?;
}
let request_body = build_request_body(request, false);
let url = format!("{}/responses", self.base_url);
let url = format!("{}/responses", self.http.base_url);
let mut req = self.build_request(&url).json(&request_body);
if let Some(t) = self.request_timeout {
if let Some(t) = self.http.request_timeout {
req = req.timeout(t);
}
let (body, headers) = send_and_read_response(req, "openai", "type").await?;
@ -972,7 +900,7 @@ impl ProviderAdapter for Adapter {
crate::provider::validate_tool_choice(self, tc)?;
}
let request_body = build_request_body(request, true);
let url = format!("{}/responses", self.base_url);
let url = format!("{}/responses", self.http.base_url);
let http_resp = self
.build_request(&url)
@ -1002,12 +930,10 @@ impl ProviderAdapter for Adapter {
let model = request.model.clone();
let rate_limit = parse_rate_limit_headers(http_resp.headers());
let byte_stream = http_resp.bytes_stream();
let stream_read_timeout = self.http.stream_read_timeout;
let stream_read_timeout = self.stream_read_timeout;
let state = SseStreamState {
byte_stream: Box::pin(byte_stream),
buffer: String::new(),
line_reader: super::common::LineReader::new(http_resp, stream_read_timeout),
model,
response_id: String::new(),
response_model: String::new(),
@ -1020,7 +946,6 @@ impl ProviderAdapter for Adapter {
emitted_text_start: false,
raw_response: None,
rate_limit,
stream_read_timeout,
};
let stream = futures::stream::unfold(state, |mut state| async move {
@ -1178,7 +1103,7 @@ mod tests {
let mut headers = HashMap::new();
headers.insert("X-Custom".to_string(), "value".to_string());
let adapter = Adapter::new("sk-test").with_default_headers(headers);
assert_eq!(adapter.default_headers.get("X-Custom").map(String::as_str), Some("value"));
assert_eq!(adapter.http.default_headers.get("X-Custom").map(String::as_str), Some("value"));
}
#[test]
@ -1186,7 +1111,7 @@ mod tests {
let adapter = Adapter::new("sk-test");
assert!(adapter.org_id.is_none());
assert!(adapter.project_id.is_none());
assert!(adapter.default_headers.is_empty());
assert!(adapter.http.default_headers.is_empty());
}
#[test]

View file

@ -18,31 +18,16 @@ use crate::types::{
/// Does NOT support reasoning tokens, built-in tools, or other Responses API
/// features. Use the primary `OpenAiAdapter` for `OpenAI`'s own API.
pub struct Adapter {
api_key: String,
base_url: String,
pub(crate) http: super::http_api::HttpApi,
provider_name: String,
default_headers: std::collections::HashMap<String, String>,
client: reqwest::Client,
request_timeout: Option<std::time::Duration>,
stream_read_timeout: Option<std::time::Duration>,
}
impl Adapter {
#[must_use]
pub fn new(api_key: impl Into<String>, base_url: 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(),
base_url: base_url.into(),
http: super::http_api::HttpApi::new(api_key, base_url),
provider_name: "openai-compatible".to_string(),
default_headers: std::collections::HashMap::new(),
client,
request_timeout: timeout.request.map(std::time::Duration::from_secs_f64),
stream_read_timeout: timeout.stream_read.map(std::time::Duration::from_secs_f64),
}
}
@ -53,30 +38,23 @@ impl Adapter {
}
#[must_use]
pub fn with_default_headers(mut self, headers: std::collections::HashMap<String, String>) -> Self {
self.default_headers = headers;
self
pub fn with_default_headers(self, headers: std::collections::HashMap<String, String>) -> Self {
Self { http: self.http.with_default_headers(headers), ..self }
}
#[must_use]
pub fn with_timeout(mut self, timeout: crate::types::AdapterTimeout) -> Self {
self.client = reqwest::Client::builder()
.connect_timeout(std::time::Duration::from_secs_f64(timeout.connect))
.build()
.unwrap_or_default();
self.request_timeout = timeout.request.map(std::time::Duration::from_secs_f64);
self.stream_read_timeout = timeout.stream_read.map(std::time::Duration::from_secs_f64);
self
pub fn with_timeout(self, timeout: crate::types::AdapterTimeout) -> Self {
Self { http: self.http.with_timeout(timeout), ..self }
}
/// Build a `reqwest::RequestBuilder` with default headers and auth.
fn build_request(&self, url: &str) -> reqwest::RequestBuilder {
let mut req = self.client.post(url);
let mut req = self.http.client.post(url);
// Apply default_headers first so adapter-specific headers can override
for (key, value) in &self.default_headers {
for (key, value) in &self.http.default_headers {
req = req.header(key, value);
}
req.bearer_auth(&self.api_key)
req.bearer_auth(&self.http.api_key)
}
}
@ -410,10 +388,10 @@ impl ProviderAdapter for Adapter {
crate::provider::validate_tool_choice(self, tc)?;
}
let api_body = build_api_request(request, None, &self.provider_name);
let url = format!("{}/chat/completions", self.base_url);
let url = format!("{}/chat/completions", self.http.base_url);
let mut req = self.build_request(&url).json(&api_body);
if let Some(t) = self.request_timeout {
if let Some(t) = self.http.request_timeout {
req = req.timeout(t);
}
let (body, headers) = send_and_read_response(req, &self.provider_name, "type").await?;
@ -482,7 +460,7 @@ impl ProviderAdapter for Adapter {
crate::provider::validate_tool_choice(self, tc)?;
}
let api_body = build_api_request(request, Some(true), &self.provider_name);
let url = format!("{}/chat/completions", self.base_url);
let url = format!("{}/chat/completions", self.http.base_url);
let http_resp = self
.build_request(&url)
@ -516,7 +494,7 @@ impl ProviderAdapter for Adapter {
let provider_name = self.provider_name.clone();
let model = request.model.clone();
let rate_limit = parse_rate_limit_headers(http_resp.headers());
let stream_read_timeout = self.stream_read_timeout;
let stream_read_timeout = self.http.stream_read_timeout;
let stream = futures::stream::unfold(
StreamState::new(http_resp, provider_name, model, rate_limit, stream_read_timeout),
@ -601,8 +579,7 @@ struct FlattenState {
/// Accumulated state while processing the SSE stream.
struct StreamState {
response: reqwest::Response,
buffer: String,
line_reader: super::common::LineReader,
provider_name: String,
model: String,
response_id: String,
@ -614,7 +591,6 @@ struct StreamState {
text_started: bool,
done: bool,
rate_limit: Option<crate::types::RateLimitInfo>,
stream_read_timeout: Option<std::time::Duration>,
}
impl StreamState {
@ -626,8 +602,7 @@ impl StreamState {
stream_read_timeout: Option<std::time::Duration>,
) -> Self {
Self {
response,
buffer: String::new(),
line_reader: super::common::LineReader::new(response, stream_read_timeout),
provider_name,
model,
response_id: String::new(),
@ -639,7 +614,6 @@ impl StreamState {
text_started: false,
done: false,
rate_limit,
stream_read_timeout,
}
}
@ -648,41 +622,11 @@ impl StreamState {
if self.done {
return Ok(None);
}
loop {
if let Some(newline_pos) = self.buffer.find('\n') {
let line = self.buffer[..newline_pos].to_string();
self.buffer = self.buffer[newline_pos + 1..].to_string();
return Ok(Some(line));
}
let chunk_result = match self.stream_read_timeout {
Some(timeout) => tokio::time::timeout(timeout, self.response.chunk()).await,
None => Ok(self.response.chunk().await),
};
match chunk_result {
Ok(Ok(Some(bytes))) => {
let text = String::from_utf8_lossy(&bytes);
self.buffer.push_str(&text);
}
Ok(Ok(None)) => {
self.done = true;
if self.buffer.is_empty() {
return Ok(None);
}
let remaining = std::mem::take(&mut self.buffer);
return Ok(Some(remaining));
}
Ok(Err(e)) => {
return Err(SdkError::Stream {
message: e.to_string(),
});
}
Err(_) => {
return Err(SdkError::Stream {
message: "stream read timed out waiting for next event".to_string(),
});
}
match self.line_reader.read_next_chunk("\n").await? {
Some(line) => Ok(Some(line)),
None => {
self.done = true;
Ok(None)
}
}
}

View file

@ -197,23 +197,11 @@ pub async fn execute_all_tools_with_repair(
async move {
let Some(t) = tool else {
return ToolResult {
tool_call_id: call_id,
content: serde_json::Value::String(format!("Unknown tool: {call_name}")),
is_error: true,
image_data: None,
image_media_type: None,
};
return ToolResult::error(call_id, format!("Unknown tool: {call_name}"));
};
let Some(handler) = &t.execute else {
return ToolResult {
tool_call_id: call_id,
content: serde_json::Value::String(format!("Unknown tool: {call_name}")),
is_error: true,
image_data: None,
image_media_type: None,
};
return ToolResult::error(call_id, format!("Unknown tool: {call_name}"));
};
let validated_args = match validate_tool_args(&args, &t.definition.parameters) {
@ -223,46 +211,24 @@ pub async fn execute_all_tools_with_repair(
match repair_fn(call_clone, validation_error).await {
Ok(repaired) => repaired,
Err(repair_error) => {
return ToolResult {
tool_call_id: call_id,
content: serde_json::Value::String(format!(
"Tool call validation failed and repair failed: {repair_error}"
)),
is_error: true,
image_data: None,
image_media_type: None,
};
return ToolResult::error(
call_id,
format!("Tool call validation failed and repair failed: {repair_error}"),
);
}
}
} else {
return ToolResult {
tool_call_id: call_id,
content: serde_json::Value::String(format!(
"Tool call validation failed: {validation_error}"
)),
is_error: true,
image_data: None,
image_media_type: None,
};
return ToolResult::error(
call_id,
format!("Tool call validation failed: {validation_error}"),
);
}
}
};
match handler(validated_args, ctx).await {
Ok(result) => ToolResult {
tool_call_id: call_id,
content: result,
is_error: false,
image_data: None,
image_media_type: None,
},
Err(err_msg) => ToolResult {
tool_call_id: call_id,
content: serde_json::Value::String(err_msg),
is_error: true,
image_data: None,
image_media_type: None,
},
Ok(result) => ToolResult::success(call_id, result),
Err(err_msg) => ToolResult::error(call_id, err_msg),
}
}
})

View file

@ -89,6 +89,28 @@ pub struct ToolResult {
pub image_media_type: Option<String>,
}
impl ToolResult {
pub fn success(id: impl Into<String>, content: serde_json::Value) -> Self {
Self {
tool_call_id: id.into(),
content,
is_error: false,
image_data: None,
image_media_type: None,
}
}
pub fn error(id: impl Into<String>, message: impl Into<String>) -> Self {
Self {
tool_call_id: id.into(),
content: serde_json::Value::String(message.into()),
is_error: true,
image_data: None,
image_media_type: None,
}
}
}
// --- 3.3 ContentPart ---
#[derive(Debug, Clone, PartialEq, Eq)]
@ -100,7 +122,6 @@ pub enum ContentPart {
ToolCall(ToolCall),
ToolResult(ToolResult),
Thinking(ThinkingData),
RedactedThinking(ThinkingData),
Other {
kind: String,
data: serde_json::Value,
@ -137,11 +158,8 @@ impl Serialize for ContentPart {
map.serialize_entry("data", v)?;
}
Self::Thinking(v) => {
map.serialize_entry("kind", "thinking")?;
map.serialize_entry("data", v)?;
}
Self::RedactedThinking(v) => {
map.serialize_entry("kind", "redacted_thinking")?;
let kind = if v.redacted { "redacted_thinking" } else { "thinking" };
map.serialize_entry("kind", kind)?;
map.serialize_entry("data", v)?;
}
Self::Other { kind, data } => {
@ -183,8 +201,8 @@ impl<'de> Deserialize<'de> for ContentPart {
"thinking" => serde_json::from_value(data)
.map(Self::Thinking)
.map_err(serde::de::Error::custom),
"redacted_thinking" => serde_json::from_value(data)
.map(Self::RedactedThinking)
"redacted_thinking" => serde_json::from_value::<ThinkingData>(data)
.map(|mut td| { td.redacted = true; Self::Thinking(td) })
.map_err(serde::de::Error::custom),
other => Ok(Self::Other {
kind: other.to_string(),
@ -733,30 +751,10 @@ pub struct GenerateResult {
pub output: Option<serde_json::Value>,
}
impl GenerateResult {
#[must_use]
pub fn text(&self) -> String {
self.response.text()
}
#[must_use]
pub fn reasoning(&self) -> Option<String> {
self.response.reasoning()
}
#[must_use]
pub fn tool_calls(&self) -> Vec<ToolCall> {
self.response.tool_calls()
}
#[must_use]
pub const fn finish_reason(&self) -> &FinishReason {
&self.response.finish_reason
}
#[must_use]
pub const fn usage(&self) -> &Usage {
&self.response.usage
impl std::ops::Deref for GenerateResult {
type Target = Response;
fn deref(&self) -> &Response {
&self.response
}
}
@ -766,35 +764,10 @@ pub struct StepResult {
pub tool_results: Vec<ToolResult>,
}
impl StepResult {
#[must_use]
pub fn text(&self) -> String {
self.response.text()
}
#[must_use]
pub fn reasoning(&self) -> Option<String> {
self.response.reasoning()
}
#[must_use]
pub fn tool_calls(&self) -> Vec<ToolCall> {
self.response.tool_calls()
}
#[must_use]
pub const fn finish_reason(&self) -> &FinishReason {
&self.response.finish_reason
}
#[must_use]
pub const fn usage(&self) -> &Usage {
&self.response.usage
}
#[must_use]
pub fn warnings(&self) -> &[Warning] {
&self.response.warnings
impl std::ops::Deref for StepResult {
type Target = Response;
fn deref(&self) -> &Response {
&self.response
}
}
@ -1244,13 +1217,7 @@ mod tests {
rate_limit: None,
};
let tool_calls = vec![ToolCall::new("call_1", "get_weather", serde_json::json!({"city": "SF"}))];
let tool_results = vec![ToolResult {
tool_call_id: "call_1".into(),
content: serde_json::json!("72F"),
is_error: false,
image_data: None,
image_media_type: None,
}];
let tool_results = vec![ToolResult::success("call_1", serde_json::json!("72F"))];
let event = StreamEvent::step_finish(
FinishReason::ToolCalls,