mirror of
https://github.com/fabro-sh/fabro.git
synced 2026-10-09 03:20:56 +00:00
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:
parent
0756a42ce8
commit
3253589193
13 changed files with 304 additions and 571 deletions
|
|
@ -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(),
|
||||
});
|
||||
|
||||
|
|
|
|||
|
|
@ -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}")),
|
||||
}
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -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)) => {
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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::*;
|
||||
|
|
|
|||
|
|
@ -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))
|
||||
}
|
||||
}
|
||||
|
||||
|
|
|
|||
54
crates/llm/src/providers/http_api.rs
Normal file
54
crates/llm/src/providers/http_api.rs
Normal 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
|
||||
}
|
||||
}
|
||||
|
|
@ -1,6 +1,7 @@
|
|||
pub mod anthropic;
|
||||
pub mod common;
|
||||
pub mod gemini;
|
||||
pub mod http_api;
|
||||
pub mod openai;
|
||||
pub mod openai_compatible;
|
||||
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
}
|
||||
}
|
||||
})
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue