diff --git a/crates/agent/src/history.rs b/crates/agent/src/history.rs index 549c990a8..36e197024 100644 --- a/crates/agent/src/history.rs +++ b/crates/agent/src/history.rs @@ -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(), }); diff --git a/crates/agent/src/session.rs b/crates/agent/src/session.rs index e957e0401..336d4e4f0 100644 --- a/crates/agent/src/session.rs +++ b/crates/agent/src/session.rs @@ -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(®istered_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}")), } } diff --git a/crates/llm/src/bin/ullm.rs b/crates/llm/src/bin/ullm.rs index aa9ffb4a7..e53bcc162 100644 --- a/crates/llm/src/bin/ullm.rs +++ b/crates/llm/src/bin/ullm.rs @@ -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)) => { diff --git a/crates/llm/src/generate.rs b/crates/llm/src/generate.rs index ff0b94282..547b6c8a9 100644 --- a/crates/llm/src/generate.rs +++ b/crates/llm/src/generate.rs @@ -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( diff --git a/crates/llm/src/providers/anthropic.rs b/crates/llm/src/providers/anthropic.rs index fcc0c9c44..399e496e1 100644 --- a/crates/llm/src/providers/anthropic.rs +++ b/crates/llm/src/providers/anthropic.rs @@ -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, - client: reqwest::Client, - request_timeout: Option, - stream_read_timeout: Option, + pub(crate) http: super::http_api::HttpApi, } impl Adapter { #[must_use] pub fn new(api_key: impl Into) -> 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) -> 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) -> Self { - self.default_headers = headers; - self + pub fn with_default_headers(self, headers: std::collections::HashMap) -> 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 { .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 { "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 { } 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>, - buffer: String, + line_reader: super::common::LineReader, accumulator: StreamAccumulator, pending_events: std::collections::VecDeque, - done: bool, /// When true, `tool_use` events for the synthetic tool are converted to text events. json_schema_mode: bool, - stream_read_timeout: Option, } impl SseReaderState { fn new( - byte_stream: impl futures::Stream> - + Send - + 'static, + http_resp: reqwest::Response, rate_limit: Option, json_schema_mode: bool, stream_read_timeout: Option, ) -> 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 { - // 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 { let mut event_type = String::new(); let mut data_parts: Vec = 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. diff --git a/crates/llm/src/providers/common.rs b/crates/llm/src/providers/common.rs index 548b81c15..06a2785c9 100644 --- a/crates/llm/src/providers/common.rs +++ b/crates/llm/src/providers/common.rs @@ -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, +} + +impl LineReader { + pub fn new(response: reqwest::Response, stream_read_timeout: Option) -> 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, 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::*; diff --git a/crates/llm/src/providers/gemini.rs b/crates/llm/src/providers/gemini.rs index 760ede6f2..b6b924ee4 100644 --- a/crates/llm/src/providers/gemini.rs +++ b/crates/llm/src/providers/gemini.rs @@ -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, - client: reqwest::Client, - request_timeout: Option, - stream_read_timeout: Option, + pub(crate) http: super::http_api::HttpApi, } impl Adapter { #[must_use] pub fn new(api_key: impl Into) -> 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) -> 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) -> Self { - self.default_headers = headers; - self + pub fn with_default_headers(self, headers: std::collections::HashMap) -> 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, /// 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, - stream_read_timeout: Option, } impl SseStreamState { fn new(http_resp: reqwest::Response, model: String, rate_limit: Option, stream_read_timeout: Option) -> 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, 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)) } } diff --git a/crates/llm/src/providers/http_api.rs b/crates/llm/src/providers/http_api.rs new file mode 100644 index 000000000..5f55d2407 --- /dev/null +++ b/crates/llm/src/providers/http_api.rs @@ -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, + pub(crate) client: reqwest::Client, + pub(crate) request_timeout: Option, + pub(crate) stream_read_timeout: Option, +} + +impl HttpApi { + #[must_use] + pub fn new(api_key: impl Into, base_url: impl Into) -> 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) -> Self { + self.default_headers = headers; + self + } +} diff --git a/crates/llm/src/providers/mod.rs b/crates/llm/src/providers/mod.rs index ff0421985..3f329589d 100644 --- a/crates/llm/src/providers/mod.rs +++ b/crates/llm/src/providers/mod.rs @@ -1,6 +1,7 @@ pub mod anthropic; pub mod common; pub mod gemini; +pub mod http_api; pub mod openai; pub mod openai_compatible; diff --git a/crates/llm/src/providers/openai.rs b/crates/llm/src/providers/openai.rs index 64d649d5d..ebe003e77 100644 --- a/crates/llm/src/providers/openai.rs +++ b/crates/llm/src/providers/openai.rs @@ -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, project_id: Option, - default_headers: std::collections::HashMap, - client: reqwest::Client, - request_timeout: Option, - stream_read_timeout: Option, } impl Adapter { #[must_use] pub fn new(api_key: impl Into) -> 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) -> 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) -> Self { - self.default_headers = headers; - self + pub fn with_default_headers(self, headers: std::collections::HashMap) -> 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, bool) { /// Mutable state carried through SSE stream processing. struct SseStreamState { - byte_stream: std::pin::Pin< - Box> + 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, rate_limit: Option, - stream_read_timeout: Option, } -/// 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)> { - 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)> { + let mut current_event: Option = 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 = 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)>, -) -> Vec { - 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, 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] diff --git a/crates/llm/src/providers/openai_compatible.rs b/crates/llm/src/providers/openai_compatible.rs index 21d0415a8..78c965d93 100644 --- a/crates/llm/src/providers/openai_compatible.rs +++ b/crates/llm/src/providers/openai_compatible.rs @@ -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, - client: reqwest::Client, - request_timeout: Option, - stream_read_timeout: Option, } impl Adapter { #[must_use] pub fn new(api_key: impl Into, base_url: impl Into) -> 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) -> Self { - self.default_headers = headers; - self + pub fn with_default_headers(self, headers: std::collections::HashMap) -> 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, - stream_read_timeout: Option, } impl StreamState { @@ -626,8 +602,7 @@ impl StreamState { stream_read_timeout: Option, ) -> 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) } } } diff --git a/crates/llm/src/tools.rs b/crates/llm/src/tools.rs index 332b3c6f0..caaf5ae71 100644 --- a/crates/llm/src/tools.rs +++ b/crates/llm/src/tools.rs @@ -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), } } }) diff --git a/crates/llm/src/types.rs b/crates/llm/src/types.rs index 1dbfd0ea1..b54f780e3 100644 --- a/crates/llm/src/types.rs +++ b/crates/llm/src/types.rs @@ -89,6 +89,28 @@ pub struct ToolResult { pub image_media_type: Option, } +impl ToolResult { + pub fn success(id: impl Into, 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, message: impl Into) -> 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::(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, } -impl GenerateResult { - #[must_use] - pub fn text(&self) -> String { - self.response.text() - } - - #[must_use] - pub fn reasoning(&self) -> Option { - self.response.reasoning() - } - - #[must_use] - pub fn tool_calls(&self) -> Vec { - 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, } -impl StepResult { - #[must_use] - pub fn text(&self) -> String { - self.response.text() - } - - #[must_use] - pub fn reasoning(&self) -> Option { - self.response.reasoning() - } - - #[must_use] - pub fn tool_calls(&self) -> Vec { - 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,