From db71e7334e48ee0fb2f66c0b47a99b2a7c74a5de Mon Sep 17 00:00:00 2001 From: Endi Sukaj Date: Sun, 30 Aug 2026 22:36:02 +0200 Subject: [PATCH 1/3] Cleanups --- src/claude/tests.rs | 2 +- src/client/mod.rs | 37 +++++++++++++------------------------ src/client/tests.rs | 30 +++++++++++++++++++++++++----- src/gemini/base.rs | 6 +----- src/gemini/tests.rs | 4 ++-- src/gemini/types.rs | 5 +++-- src/openai/tests.rs | 2 +- 7 files changed, 46 insertions(+), 40 deletions(-) diff --git a/src/claude/tests.rs b/src/claude/tests.rs index d803c02..8d88377 100644 --- a/src/claude/tests.rs +++ b/src/claude/tests.rs @@ -22,7 +22,7 @@ fn make_model(model: ClaudeModel) -> ClaudeApiModel { fn default_settings() -> Settings { Settings { - max_tokens: Some(8000), + max_tokens: Some(8000.0), timeout: None, temperature: None, thinking_budget: None, diff --git a/src/client/mod.rs b/src/client/mod.rs index 9797182..ddabbc8 100644 --- a/src/client/mod.rs +++ b/src/client/mod.rs @@ -47,8 +47,9 @@ pub enum Role { User, } -#[derive(Debug, Clone, PartialEq, Eq)] +#[derive(Debug, Clone, Default, PartialEq, Eq, Serialize, Deserialize)] pub enum MessageType { + #[default] Text, FunctionCall(FunctionCall), FunctionResponse { @@ -57,17 +58,11 @@ pub enum MessageType { }, } -impl Default for MessageType { - fn default() -> Self { - MessageType::Text - } -} - #[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] pub struct Message { pub content: String, pub role: Option, - #[serde(skip)] + #[serde(default)] pub message_type: MessageType, } @@ -113,7 +108,7 @@ impl Message { } #[async_trait] -pub trait Model { +pub trait Model: Send + Sync { async fn completion( &self, request: ModelRequest, @@ -136,9 +131,9 @@ pub trait Model { #[derive(Clone)] pub struct Settings { - pub max_tokens: Option, + pub max_tokens: Option, pub timeout: Option, - pub temperature: Option, + pub temperature: Option, pub thinking_budget: Option, } @@ -180,14 +175,11 @@ impl Tool { let arg_schema = schema_for!(T); let json_value = serde_json::to_value(&arg_schema)?; let parameters: ToolParameters = serde_json::from_value(json_value)?; - match self.parameters { - None => Ok(Tool { - name: self.name, - description: self.description, - parameters: Some(parameters), - }), - Some(_) => Ok(self), - } + Ok(Tool { + name: self.name, + description: self.description, + parameters: Some(parameters), + }) } } @@ -200,9 +192,6 @@ pub struct ModelRequestBuilder<'a> { pub tools: Option>, } -unsafe impl<'a> Sync for ModelRequestBuilder<'a> {} -unsafe impl<'a> Send for ModelRequestBuilder<'a> {} - pub struct ModelRequest { pub system: Option, pub messages: Option>, @@ -251,7 +240,7 @@ impl<'a> ModelRequestBuilder<'a> { match self.tools { None => self.tools = Some(vec![tool]), Some(_) => { - self.tools.clone().map(|mut ts| ts.push(tool)); + self.tools.get_or_insert_with(Vec::new).push(tool); } } return self; @@ -261,7 +250,7 @@ impl<'a> ModelRequestBuilder<'a> { match self.tools { None => self.tools = Some(tools), Some(_) => { - self.tools.clone().map(|mut ts| ts.extend(tools)); + self.tools.get_or_insert_with(Vec::new).extend(tools); } } return self; diff --git a/src/client/tests.rs b/src/client/tests.rs index 1ae0aa0..8885ed1 100644 --- a/src/client/tests.rs +++ b/src/client/tests.rs @@ -123,17 +123,17 @@ fn test_with_settings() { let model = MockModel; let mut builder = ModelRequestBuilder::new(&model); let settings = Settings { - max_tokens: Some(100), + max_tokens: Some(100.0), timeout: Some(30), - temperature: Some(7), + temperature: Some(0.7), thinking_budget: None, }; builder.with_settings(settings); let s = builder.settings.unwrap(); - assert_eq!(s.max_tokens, Some(100)); + assert_eq!(s.max_tokens, Some(100.0)); assert_eq!(s.timeout, Some(30)); - assert_eq!(s.temperature, Some(7)); + assert_eq!(s.temperature, Some(0.7)); } #[test] @@ -186,7 +186,7 @@ fn test_chaining() { .with_message(Message::user("User msg".to_string())) .with_tool(Tool::new("tool", "desc")) .with_settings(Settings { - max_tokens: Some(50), + max_tokens: Some(50.0), timeout: None, temperature: None, thinking_budget: None, @@ -273,3 +273,23 @@ fn test_model_name() { let model_name = model.model_name(); assert_eq!(model_name, "test-model".to_string()); } + +#[test] +fn test_message_serde_roundtrip_preserves_message_type() { + let fc = FunctionCall { + name: "get_weather".to_string(), + args: HashMap::from([("city".to_string(), serde_json::json!("Berlin"))]), + }; + let original = Message::function_call(fc); + + let json = serde_json::to_string(&original).unwrap(); + let back: Message = serde_json::from_str(&json).unwrap(); + + assert_eq!(back, original); + assert!(matches!(back.message_type, MessageType::FunctionCall(_))); + + // Messages serialized before `message_type` existed still deserialize (as Text). + let legacy = r#"{"content":"hi","role":"user"}"#; + let msg: Message = serde_json::from_str(legacy).unwrap(); + assert_eq!(msg.message_type, MessageType::Text); +} diff --git a/src/gemini/base.rs b/src/gemini/base.rs index 39c1427..7a3a566 100644 --- a/src/gemini/base.rs +++ b/src/gemini/base.rs @@ -25,11 +25,7 @@ pub trait GeminiClient: Model { let generation_config = GenerationConfig { max_output_tokens: request.settings.clone().and_then(|s| s.max_tokens), - temperature: request - .settings - .clone() - .map(|s| s.temperature.unwrap_or_default()) - .unwrap_or_default(), + temperature: request.settings.clone().and_then(|s| s.temperature), thinking_config, }; diff --git a/src/gemini/tests.rs b/src/gemini/tests.rs index 78c0cd9..b0b052c 100644 --- a/src/gemini/tests.rs +++ b/src/gemini/tests.rs @@ -32,7 +32,7 @@ fn make_vertex(model: GeminiModel) -> GeminiVertexModel { fn default_settings() -> Settings { Settings { - max_tokens: Some(8000), + max_tokens: Some(8000.0), timeout: None, temperature: None, // Use dynamic thinking (-1) so thinking-only models like Gemini 3.1 Pro @@ -443,7 +443,7 @@ fn request_with_thinking(thinking_budget: Option) -> crate::client::ModelRe system: None, messages: Some(vec![Message::user("hi".to_string())]), settings: Some(Settings { - max_tokens: Some(100), + max_tokens: Some(100.0), timeout: None, temperature: None, thinking_budget, diff --git a/src/gemini/types.rs b/src/gemini/types.rs index 83972e7..f2cdcf4 100644 --- a/src/gemini/types.rs +++ b/src/gemini/types.rs @@ -32,8 +32,9 @@ pub struct ThinkingConfig { #[derive(Serialize)] pub struct GenerationConfig { #[serde(rename = "maxOutputTokens", skip_serializing_if = "Option::is_none")] - pub max_output_tokens: Option, - pub temperature: i16, + pub max_output_tokens: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub temperature: Option, #[serde(rename = "thinkingConfig", skip_serializing_if = "Option::is_none")] pub thinking_config: Option, } diff --git a/src/openai/tests.rs b/src/openai/tests.rs index 1809744..d73cd3e 100644 --- a/src/openai/tests.rs +++ b/src/openai/tests.rs @@ -22,7 +22,7 @@ fn make_model(model: OpenAiModel) -> OpenAiApiModel { fn default_settings() -> Settings { Settings { - max_tokens: Some(8000), + max_tokens: Some(8000.0), timeout: None, temperature: None, thinking_budget: None, From e108cad64cdb41df8fb9fd19b90fc49f0f22fa79 Mon Sep 17 00:00:00 2001 From: Endi Sukaj Date: Sun, 30 Aug 2026 22:54:59 +0200 Subject: [PATCH 2/3] Invert provider abstraction --- README.md | 54 ++--- src/claude/adapter.rs | 213 +++++++++++++++++++ src/claude/base.rs | 314 ----------------------------- src/claude/direct_api_client.rs | 59 ------ src/claude/mod.rs | 57 +++++- src/claude/tests.rs | 27 +-- src/gemini/adapter.rs | 139 +++++++++++++ src/gemini/base.rs | 224 -------------------- src/gemini/direct_api_client.rs | 61 ------ src/gemini/mod.rs | 123 ++++++++++- src/gemini/tests.rs | 54 ++--- src/gemini/vertex_client.rs | 63 ------ src/lib.rs | 2 + src/openai/{base.rs => adapter.rs} | 155 +++----------- src/openai/direct_api_client.rs | 58 ------ src/openai/mod.rs | 55 ++++- src/openai/tests.rs | 39 +--- src/provider.rs | 170 ++++++++++++++++ 18 files changed, 838 insertions(+), 1029 deletions(-) create mode 100644 src/claude/adapter.rs delete mode 100644 src/claude/base.rs delete mode 100644 src/claude/direct_api_client.rs create mode 100644 src/gemini/adapter.rs delete mode 100644 src/gemini/base.rs delete mode 100644 src/gemini/direct_api_client.rs delete mode 100644 src/gemini/vertex_client.rs rename src/openai/{base.rs => adapter.rs} (51%) delete mode 100644 src/openai/direct_api_client.rs create mode 100644 src/provider.rs diff --git a/README.md b/README.md index e99ab13..70a717b 100644 --- a/README.md +++ b/README.md @@ -62,8 +62,6 @@ futures = "0.3" schemars = "1" # only needed if you define tools serde = { version = "1", features = ["derive"] } ``` -de = { version = "1", features = ["derive"] } -``` Environment variables used by the examples: @@ -83,11 +81,7 @@ use langrust::{ClaudeApiModel, ClaudeModel, Message, Model}; #[tokio::main] async fn main() -> Result<(), Box> { - let model = ClaudeApiModel { - api_key: std::env::var("ANTHROPIC_API_KEY")?, - client: reqwest::Client::new(), - model: ClaudeModel::Sonnet4_5, - }; + let model = ClaudeApiModel::new(std::env::var("ANTHROPIC_API_KEY")?, ClaudeModel::Sonnet4_5); let completion = model .new_request() @@ -111,11 +105,7 @@ use langrust::{Message, Model, OpenAiApiModel, OpenAiModel, Settings}; #[tokio::main] async fn main() -> Result<(), Box> { - let model = OpenAiApiModel { - api_key: std::env::var("OPENAI_API_KEY")?, - client: reqwest::Client::new(), - model: OpenAiModel::Gpt5_4Mini, - }; + let model = OpenAiApiModel::new(std::env::var("OPENAI_API_KEY")?, OpenAiModel::Gpt5_4Mini); let settings = Settings { max_tokens: Some(256), @@ -144,11 +134,7 @@ use langrust::{GeminiApiModel, GeminiModel, Message, Model}; #[tokio::main] async fn main() -> Result<(), Box> { - let model = GeminiApiModel { - api_key: std::env::var("GEMINI_KEY")?, - client: reqwest::Client::new(), - model: GeminiModel::Gemini25Flash, - }; + let model = GeminiApiModel::new(std::env::var("GEMINI_KEY")?, GeminiModel::Gemini25Flash); let history = vec![ Message::user("What is the capital of Japan?".to_string()), @@ -177,11 +163,7 @@ use langrust::{GeminiModel, GeminiVertexModel, Message, Model}; #[tokio::main] async fn main() -> Result<(), Box> { - let model = GeminiVertexModel { - project_name: std::env::var("VERTEX_PROJECT")?, - client: reqwest::Client::new(), - model: GeminiModel::Gemini31Pro, - }; + let model = GeminiVertexModel::new(std::env::var("VERTEX_PROJECT")?, GeminiModel::Gemini31Pro); let completion = model .new_request() @@ -205,11 +187,7 @@ use langrust::{ClaudeApiModel, ClaudeModel, Message, Model, StreamEvent}; #[tokio::main] async fn main() -> Result<(), Box> { - let model = ClaudeApiModel { - api_key: std::env::var("ANTHROPIC_API_KEY")?, - client: reqwest::Client::new(), - model: ClaudeModel::Sonnet4_5, - }; + let model = ClaudeApiModel::new(std::env::var("ANTHROPIC_API_KEY")?, ClaudeModel::Sonnet4_5); let mut stream = model .new_request() @@ -250,11 +228,7 @@ struct GetWeatherArgs { #[tokio::main] async fn main() -> Result<(), Box> { - let model = OpenAiApiModel { - api_key: std::env::var("OPENAI_API_KEY")?, - client: reqwest::Client::new(), - model: OpenAiModel::Gpt5_4, - }; + let model = OpenAiApiModel::new(std::env::var("OPENAI_API_KEY")?, OpenAiModel::Gpt5_4); let tool = Tool::new("get_weather", "Fetch the current weather for a city.") .with_parameter::()?; @@ -295,6 +269,22 @@ The same `Tool` value can be passed to `ClaudeApiModel`, `GeminiApiModel` or `GeminiVertexModel` — the schema is translated automatically (including Gemini's uppercase type names and nullable-handling). +## Architecture + +Every client is the generic `LlmClient` composed from two small traits: + +- `ProviderAdapter` (`A`) — pure request/response conversions for one wire + format (no I/O): `ClaudeAdapter`, `OpenAiAdapter`, `GeminiAdapter`. +- `Transport` (`T`) — endpoint + authentication (only I/O): e.g. + `AnthropicTransport`, `OpenAiTransport`, `GeminiApiTransport`, `VertexTransport`. + +`Model` is implemented once for `LlmClient`, and the public client types are +aliases such as `type ClaudeApiModel = LlmClient`. +Direct Gemini and Vertex AI share `GeminiAdapter` and differ only in transport. +To use a custom `reqwest::Client`, call the `with_client` constructors; to add +a new backend, implement `Transport` (new auth/endpoint for an existing wire +format) or `ProviderAdapter` + `Transport` (entirely new provider). + ## Core types cheat-sheet - `Model` — trait with `completion()` and `stream_completion()`; all providers implement it. diff --git a/src/claude/adapter.rs b/src/claude/adapter.rs new file mode 100644 index 0000000..d10fd2d --- /dev/null +++ b/src/claude/adapter.rs @@ -0,0 +1,213 @@ +use std::collections::HashMap; + +use crate::{ + claude::types::{ + BlockDelta, ClaudeMessage, ClaudeRequest, ClaudeResponse, ClaudeTool, ContentBlock, + DEFAULT_MAX_TOKENS, ResponseBlock, StreamContentBlock, StreamingEvent, ThinkingConfig, + synth_tool_use_id, + }, + client::{Completion, FunctionCall, MessageType, ModelRequest, StreamEvent, Usage}, + provider::{BoxError, ProviderAdapter}, +}; + +/// Pure conversions between the common request/response types and the +/// Anthropic Messages API wire format. No I/O. +#[derive(Debug, Default, Clone, Copy)] +pub struct ClaudeAdapter; + +impl ProviderAdapter for ClaudeAdapter { + const NAME: &'static str = "Claude"; + type Request = ClaudeRequest; + type StreamState = ClaudeStreamState; + + fn build_body(&self, request: &ModelRequest, model: &str, stream: bool) -> ClaudeRequest { + let settings = request.settings.as_ref(); + + let max_tokens = settings + .and_then(|s| s.max_tokens) + .map(|v| v as i32) + .unwrap_or(DEFAULT_MAX_TOKENS); + + let temperature = settings.and_then(|s| s.temperature); + + let thinking = settings + .and_then(|s| s.thinking_budget) + .filter(|b| *b > 0) + .map(|b| ThinkingConfig { + kind: "enabled", + budget_tokens: b as i32, + }); + + let messages: Vec = request + .messages + .clone() + .unwrap_or_default() + .iter() + .map(|m| match &m.message_type { + MessageType::Text => ClaudeMessage { + role: match m.role { + Some(crate::client::Role::Model) => "assistant", + _ => "user", + }, + content: vec![ContentBlock::Text { + text: m.content.clone(), + }], + }, + MessageType::FunctionCall(fc) => ClaudeMessage { + role: "assistant", + content: vec![ContentBlock::ToolUse { + id: synth_tool_use_id(&fc.name), + name: fc.name.clone(), + input: fc.args.clone(), + }], + }, + MessageType::FunctionResponse { name, response } => ClaudeMessage { + role: "user", + content: vec![ContentBlock::ToolResult { + tool_use_id: synth_tool_use_id(name), + content: response + .as_ref() + .map(|v| v.to_string()) + .unwrap_or_else(|| "null".to_string()), + }], + }, + }) + .collect(); + + let tools = request + .tools + .as_ref() + .map(|ts| ts.iter().map(ClaudeTool::from_tool).collect()); + + ClaudeRequest { + model: model.to_string(), + max_tokens, + system: request.system.clone(), + messages, + temperature, + tools, + thinking, + stream: if stream { Some(true) } else { None }, + } + } + + fn parse_completion(&self, body: &[u8]) -> Result { + let body: ClaudeResponse = serde_json::from_slice(body)?; + + let mut text = String::new(); + let mut function: Option = None; + for block in body.content { + match block { + ResponseBlock::Text { text: t } => text.push_str(&t), + ResponseBlock::ToolUse { name, input, .. } => { + function = Some(FunctionCall { name, args: input }); + } + ResponseBlock::Other => {} + } + } + + let total = body.usage.input_tokens + body.usage.output_tokens; + Ok(Completion { + completion: text, + usage: Usage { + prompt_tokens: body.usage.input_tokens, + completion_tokens: body.usage.output_tokens, + total_tokens: total, + }, + function, + }) + } + + fn map_sse_event(&self, data: &str, state: &mut ClaudeStreamState) -> Vec { + if data == "[DONE]" { + return Vec::new(); + } + let mut out = Vec::new(); + match serde_json::from_str::(data) { + Err(e) => out.push(StreamEvent::Error(e.to_string())), + Ok(ev) => handle_event(ev, state, &mut out), + } + out + } +} + +#[derive(Default)] +pub struct ClaudeStreamState { + tool_blocks: HashMap, + prompt_tokens: i32, +} + +struct ToolBlockAcc { + name: String, + json_buf: String, +} + +fn handle_event(ev: StreamingEvent, state: &mut ClaudeStreamState, out: &mut Vec) { + match ev { + StreamingEvent::MessageStart { message } => { + state.prompt_tokens = message.usage.input_tokens; + } + StreamingEvent::ContentBlockStart { + index, + content_block, + } => { + if let StreamContentBlock::ToolUse { name, .. } = content_block { + state.tool_blocks.insert( + index, + ToolBlockAcc { + name, + json_buf: String::new(), + }, + ); + } + } + StreamingEvent::ContentBlockDelta { index, delta } => match delta { + BlockDelta::TextDelta { text } => { + if !text.is_empty() { + out.push(StreamEvent::Delta(text)); + } + } + BlockDelta::InputJsonDelta { partial_json } => { + if let Some(acc) = state.tool_blocks.get_mut(&index) { + acc.json_buf.push_str(&partial_json); + } + } + BlockDelta::Other => {} + }, + StreamingEvent::ContentBlockStop { index } => { + if let Some(acc) = state.tool_blocks.remove(&index) { + let args: HashMap = if acc.json_buf.is_empty() { + HashMap::new() + } else { + match serde_json::from_str(&acc.json_buf) { + Ok(v) => v, + Err(e) => { + out.push(StreamEvent::Error(format!( + "failed to parse streamed tool input JSON: {}", + e + ))); + return; + } + } + }; + out.push(StreamEvent::FunctionCall(FunctionCall { + name: acc.name, + args, + })); + } + } + StreamingEvent::MessageDelta { usage, .. } => { + out.push(StreamEvent::Usage(Usage { + prompt_tokens: state.prompt_tokens, + completion_tokens: usage.output_tokens, + total_tokens: state.prompt_tokens + usage.output_tokens, + })); + } + StreamingEvent::MessageStop => {} + StreamingEvent::Ping => {} + StreamingEvent::Error { error } => { + out.push(StreamEvent::Error(error.message)); + } + StreamingEvent::Other => {} + } +} diff --git a/src/claude/base.rs b/src/claude/base.rs deleted file mode 100644 index 7113c6e..0000000 --- a/src/claude/base.rs +++ /dev/null @@ -1,314 +0,0 @@ -use std::collections::HashMap; -use std::error::Error; - -use eventsource_stream::Eventsource; -use futures::{StreamExt, TryFutureExt, stream}; -use reqwest::RequestBuilder; - -use crate::{ - claude::types::{ - BlockDelta, ClaudeMessage, ClaudeRequest, ClaudeResponse, ClaudeTool, ContentBlock, - DEFAULT_MAX_TOKENS, ResponseBlock, StreamContentBlock, StreamingEvent, ThinkingConfig, - synth_tool_use_id, - }, - client::{ - Completion, FunctionCall, MessageType, Model, ModelRequest, StreamEvent, StreamResult, - Usage, - }, -}; - -pub trait ClaudeClient: Model { - fn create_request_body(&self, request: ModelRequest, stream: bool) -> ClaudeRequest { - let settings = request.settings.clone(); - - let max_tokens = settings - .as_ref() - .and_then(|s| s.max_tokens) - .map(|v| v as i32) - .unwrap_or(DEFAULT_MAX_TOKENS); - - let temperature = settings - .as_ref() - .and_then(|s| s.temperature) - .map(|v| v as f32); - - // Extended thinking: enabled iff caller passed a non-zero budget. - let thinking = settings - .as_ref() - .and_then(|s| s.thinking_budget) - .filter(|b| *b > 0) - .map(|b| ThinkingConfig { - kind: "enabled", - budget_tokens: b as i32, - }); - - let messages: Vec = request - .messages - .clone() - .unwrap_or_default() - .iter() - .map(|m| match &m.message_type { - MessageType::Text => ClaudeMessage { - role: match m.role { - Some(crate::client::Role::Model) => "assistant", - _ => "user", - }, - content: vec![ContentBlock::Text { - text: m.content.clone(), - }], - }, - MessageType::FunctionCall(fc) => ClaudeMessage { - role: "assistant", - content: vec![ContentBlock::ToolUse { - id: synth_tool_use_id(&fc.name), - name: fc.name.clone(), - input: fc.args.clone(), - }], - }, - MessageType::FunctionResponse { name, response } => ClaudeMessage { - role: "user", - content: vec![ContentBlock::ToolResult { - tool_use_id: synth_tool_use_id(name), - content: response - .as_ref() - .map(|v| v.to_string()) - .unwrap_or_else(|| "null".to_string()), - }], - }, - }) - .collect(); - - let tools = request - .tools - .clone() - .map(|ts| ts.iter().map(ClaudeTool::from_tool).collect()); - - ClaudeRequest { - model: self.model_name(), - max_tokens, - system: request.system.clone(), - messages, - temperature, - tools, - thinking, - stream: if stream { Some(true) } else { None }, - } - } - - async fn generate_content( - &self, - request: ModelRequest, - ) -> Result> { - let endpoint = self.get_endpoint(); - let body = self.create_request_body(request, false); - let response = self.build_request(&endpoint, &body).await?.send().await?; - - let status = response.status(); - if !status.is_success() { - let err = response.text().map_err(|e| e.to_string()).await?; - return Err(format!("Claude request failed with status {}: {}", status, err).into()); - } - - let body: ClaudeResponse = response.json().await?; - - let mut text = String::new(); - let mut function: Option = None; - for block in body.content { - match block { - ResponseBlock::Text { text: t } => text.push_str(&t), - ResponseBlock::ToolUse { name, input, .. } => { - function = Some(FunctionCall { name, args: input }); - } - ResponseBlock::Other => {} - } - } - - let total = body.usage.input_tokens + body.usage.output_tokens; - Ok(Completion { - completion: text, - usage: Usage { - prompt_tokens: body.usage.input_tokens, - completion_tokens: body.usage.output_tokens, - total_tokens: total, - }, - function, - }) - } - - async fn stream_generate_content( - &self, - request: ModelRequest, - ) -> Result> { - let endpoint = self.get_endpoint(); - let body = self.create_request_body(request, true); - let response = self.build_request(&endpoint, &body).await?.send().await?; - - let status = response.status(); - if !status.is_success() { - let err = response.text().map_err(|e| e.to_string()).await?; - return Err(format!( - "Claude streaming request failed with status {}: {}", - status, err - ) - .into()); - } - - // State threaded through `unfold`. Defined at module scope below. - let sse = Box::pin(response.bytes_stream().eventsource()); - let state = State { - sse, - buffer: std::collections::VecDeque::new(), - tool_blocks: HashMap::new(), - prompt_tokens: 0, - }; - - let out = stream::unfold(state, |mut state| async move { - loop { - if let Some(ev) = state.buffer.pop_front() { - return Some((ev, state)); - } - - let next = state.sse.next().await?; - match next { - Err(e) => { - state.buffer.push_back(StreamEvent::Error(e.to_string())); - } - Ok(event) => { - if event.data.is_empty() || event.data == "[DONE]" { - continue; - } - let parsed: Result = serde_json::from_str(&event.data); - match parsed { - Err(e) => state.buffer.push_back(StreamEvent::Error(e.to_string())), - Ok(ev) => handle_event(ev, &mut state), - } - } - } - } - }); - - Ok(Box::pin(out)) - } - - fn get_endpoint(&self) -> String; - - async fn build_request( - &self, - endpoint: &String, - request_body: &ClaudeRequest, - ) -> Result>; -} - -// Separate free function so we can mutate `State` fields cleanly. -fn handle_event(ev: StreamingEvent, state: &mut State) { - match ev { - StreamingEvent::MessageStart { message } => { - state.set_prompt_tokens(message.usage.input_tokens); - } - StreamingEvent::ContentBlockStart { - index, - content_block, - } => { - if let StreamContentBlock::ToolUse { name, .. } = content_block { - state.tool_block_insert(index, name); - } - } - StreamingEvent::ContentBlockDelta { index, delta } => match delta { - BlockDelta::TextDelta { text } => { - if !text.is_empty() { - state.push_event(StreamEvent::Delta(text)); - } - } - BlockDelta::InputJsonDelta { partial_json } => { - state.tool_block_append(index, &partial_json); - } - BlockDelta::Other => {} - }, - StreamingEvent::ContentBlockStop { index } => { - if let Some((name, json_buf)) = state.tool_block_take(index) { - let args: HashMap = if json_buf.is_empty() { - HashMap::new() - } else { - match serde_json::from_str(&json_buf) { - Ok(v) => v, - Err(e) => { - state.push_event(StreamEvent::Error(format!( - "failed to parse streamed tool input JSON: {}", - e - ))); - return; - } - } - }; - state.push_event(StreamEvent::FunctionCall(FunctionCall { name, args })); - } - } - StreamingEvent::MessageDelta { usage, .. } => { - let prompt = state.prompt_tokens(); - state.push_event(StreamEvent::Usage(Usage { - prompt_tokens: prompt, - completion_tokens: usage.output_tokens, - total_tokens: prompt + usage.output_tokens, - })); - } - StreamingEvent::MessageStop => {} - StreamingEvent::Ping => {} - StreamingEvent::Error { error } => { - state.push_event(StreamEvent::Error(error.message)); - } - StreamingEvent::Other => {} - } -} - -// Streaming state (module-scope so `handle_event` can reference it). -struct State { - sse: std::pin::Pin< - Box< - dyn futures::Stream< - Item = Result< - eventsource_stream::Event, - eventsource_stream::EventStreamError, - >, - > + Send, - >, - >, - buffer: std::collections::VecDeque, - tool_blocks: HashMap, - prompt_tokens: i32, -} - -struct ToolBlockAcc { - name: String, - json_buf: String, -} - -impl State { - fn push_event(&mut self, ev: StreamEvent) { - self.buffer.push_back(ev); - } - fn set_prompt_tokens(&mut self, v: i32) { - self.prompt_tokens = v; - } - fn prompt_tokens(&self) -> i32 { - self.prompt_tokens - } - fn tool_block_insert(&mut self, index: u32, name: String) { - self.tool_blocks.insert( - index, - ToolBlockAcc { - name, - json_buf: String::new(), - }, - ); - } - fn tool_block_append(&mut self, index: u32, s: &str) { - if let Some(acc) = self.tool_blocks.get_mut(&index) { - acc.json_buf.push_str(s); - } - } - fn tool_block_take(&mut self, index: u32) -> Option<(String, String)> { - self.tool_blocks - .remove(&index) - .map(|acc| (acc.name, acc.json_buf)) - } -} diff --git a/src/claude/direct_api_client.rs b/src/claude/direct_api_client.rs deleted file mode 100644 index 46c6a54..0000000 --- a/src/claude/direct_api_client.rs +++ /dev/null @@ -1,59 +0,0 @@ -use std::error::Error; - -use async_trait::async_trait; -use reqwest::RequestBuilder; - -use crate::{ - claude::{ - base::ClaudeClient, - types::{ClaudeModel, ClaudeRequest}, - }, - client::{Completion, Model, ModelRequest, StreamResult}, -}; - -pub struct ClaudeApiModel { - pub api_key: String, - pub client: reqwest::Client, - pub model: ClaudeModel, -} - -#[async_trait] -impl Model for ClaudeApiModel { - async fn completion( - &self, - request: ModelRequest, - ) -> Result> { - self.generate_content(request).await - } - - async fn stream_completion( - &self, - request: ModelRequest, - ) -> Result> { - self.stream_generate_content(request).await - } - - fn model_name(&self) -> String { - self.model.to_string() - } -} - -impl ClaudeClient for ClaudeApiModel { - fn get_endpoint(&self) -> String { - "https://api.anthropic.com/v1/messages".to_string() - } - - async fn build_request( - &self, - endpoint: &String, - request_body: &ClaudeRequest, - ) -> Result> { - Ok(self - .client - .post(endpoint) - .header("x-api-key", self.api_key.clone()) - .header("anthropic-version", "2023-06-01") - .header("Content-Type", "application/json") - .json(request_body)) - } -} diff --git a/src/claude/mod.rs b/src/claude/mod.rs index 47da826..857af07 100644 --- a/src/claude/mod.rs +++ b/src/claude/mod.rs @@ -1,9 +1,60 @@ -mod base; -mod direct_api_client; +mod adapter; mod types; #[cfg(test)] mod tests; -pub use direct_api_client::ClaudeApiModel; +use async_trait::async_trait; + +use crate::provider::{Action, BoxError, LlmClient, Transport}; + +pub use adapter::ClaudeAdapter; pub use types::ClaudeModel; + +/// Anthropic Messages API client (API-key auth). +pub type ClaudeApiModel = LlmClient; + +impl ClaudeApiModel { + pub fn new(api_key: impl Into, model: ClaudeModel) -> Self { + Self::with_client(api_key, model, reqwest::Client::new()) + } + + pub fn with_client( + api_key: impl Into, + model: ClaudeModel, + client: reqwest::Client, + ) -> Self { + LlmClient::from_parts( + AnthropicTransport { + api_key: api_key.into(), + client, + }, + model.to_string(), + ) + } +} + +pub struct AnthropicTransport { + pub api_key: String, + pub client: reqwest::Client, +} + +#[async_trait] +impl Transport for AnthropicTransport { + async fn send( + &self, + _model: &str, + _action: Action, + body: serde_json::Value, + ) -> Result { + Ok(self + .client + .post("https://api.anthropic.com/v1/messages") + .header("x-api-key", self.api_key.clone()) + .header("anthropic-version", "2023-06-01") + .header("Content-Type", "application/json") + .json(&body) + .send() + .await?) + } +} diff --git a/src/claude/tests.rs b/src/claude/tests.rs index 8d88377..fa526ac 100644 --- a/src/claude/tests.rs +++ b/src/claude/tests.rs @@ -6,18 +6,17 @@ use serde::{Deserialize, Serialize}; use crate::{ claude::{ - direct_api_client::ClaudeApiModel, + ClaudeApiModel, types::{ClaudeModel, ClaudeTool}, }, client::{Message, Model, Settings, StreamEvent, Tool, Usage}, }; fn make_model(model: ClaudeModel) -> ClaudeApiModel { - ClaudeApiModel { - client: reqwest::Client::new(), - api_key: env::var("CLAUDE_KEY").expect("CLAUDE_KEY env var must be set"), + ClaudeApiModel::new( + env::var("CLAUDE_KEY").expect("CLAUDE_KEY env var must be set"), model, - } + ) } fn default_settings() -> Settings { @@ -345,24 +344,12 @@ fn test_claude_tool_emits_required_and_name_and_description() { #[test] fn test_model_name_claude_api() { - let m = ClaudeApiModel { - client: reqwest::Client::new(), - api_key: "dummy-key".to_string(), - model: ClaudeModel::Sonnet4_5, - }; + let m = ClaudeApiModel::new("dummy-key", ClaudeModel::Sonnet4_5); assert_eq!(m.model_name(), "claude-sonnet-4-5"); - let m = ClaudeApiModel { - client: reqwest::Client::new(), - api_key: "dummy-key".to_string(), - model: ClaudeModel::Opus4_6, - }; + let m = ClaudeApiModel::new("dummy-key", ClaudeModel::Opus4_6); assert_eq!(m.model_name(), "claude-opus-4-6"); - let m = ClaudeApiModel { - client: reqwest::Client::new(), - api_key: "dummy-key".to_string(), - model: ClaudeModel::Opus4_7, - }; + let m = ClaudeApiModel::new("dummy-key", ClaudeModel::Opus4_7); assert_eq!(m.model_name(), "claude-opus-4-7"); } diff --git a/src/gemini/adapter.rs b/src/gemini/adapter.rs new file mode 100644 index 0000000..fe3cef8 --- /dev/null +++ b/src/gemini/adapter.rs @@ -0,0 +1,139 @@ +use crate::{ + client::{Completion, FunctionCall, MessageType, ModelRequest, Role, StreamEvent, Usage}, + gemini::types::{ + Content, FunctionCallPart, FunctionResponsePart, GeminiRequest, GeminiResponse, GeminiTool, + GeminiTools, GenerationConfig, Part, SystemInstructionContent, ThinkingConfig, + }, + provider::{BoxError, ProviderAdapter}, +}; + +#[derive(Debug, Default, Clone, Copy)] +pub struct GeminiAdapter; + +impl ProviderAdapter for GeminiAdapter { + const NAME: &'static str = "Gemini"; + type Request = GeminiRequest; + type StreamState = (); + + fn build_body(&self, request: &ModelRequest, _model: &str, _stream: bool) -> GeminiRequest { + // Model and streaming are part of the endpoint URL, not the body. + let settings = request.settings.as_ref(); + + let thinking_config = settings + .and_then(|s| s.thinking_budget) + .map(|thinking_budget| ThinkingConfig { thinking_budget }); + + let generation_config = GenerationConfig { + max_output_tokens: settings.and_then(|s| s.max_tokens), + temperature: settings.and_then(|s| s.temperature), + thinking_config, + }; + + let contents: Vec = request + .messages + .clone() + .unwrap_or_default() + .iter() + .map(|message| match &message.message_type { + MessageType::Text => Content { + parts: vec![Part::Text { + text: message.content.clone(), + }], + role: message.role.clone().unwrap_or(Role::User), + }, + MessageType::FunctionCall(fc) => Content { + parts: vec![Part::FunctionCall { + function_call: FunctionCallPart { + name: fc.name.clone(), + args: fc.args.clone(), + }, + }], + role: Role::Model, + }, + MessageType::FunctionResponse { name, response } => Content { + parts: vec![Part::FunctionResponse { + function_response: FunctionResponsePart { + name: name.clone(), + response: response.clone().unwrap_or(serde_json::Value::Null), + }, + }], + role: Role::User, + }, + }) + .collect(); + + let system_instruction = request.system.clone().map(|m| SystemInstructionContent { + parts: vec![Part::Text { text: m }], + }); + + GeminiRequest { + system_instruction, + contents, + generation_config, + tools: request.tools.as_ref().map(|ts| { + vec![GeminiTools { + function_declarations: ts.iter().map(GeminiTool::from_tool).collect(), + }] + }), + } + } + + fn parse_completion(&self, body: &[u8]) -> Result { + let response_body: GeminiResponse = serde_json::from_slice(body)?; + + let content = response_body + .get_text() + .ok_or_else(|| -> BoxError { "Missing completion from response".into() })?; + + Ok(Completion { + completion: content, + usage: Usage { + prompt_tokens: response_body.get_prompt_tokens().unwrap_or(0), + completion_tokens: response_body.get_completion_tokens().unwrap_or(0), + total_tokens: response_body.get_total_tokens().unwrap_or(0), + }, + function: response_body.get_function().map(|gf| FunctionCall { + name: gf.name, + args: gf.args, + }), + }) + } + + fn map_sse_event(&self, data: &str, _state: &mut ()) -> Vec { + let response: GeminiResponse = match serde_json::from_str(data) { + Ok(r) => r, + Err(e) => return vec![StreamEvent::Error(e.to_string())], + }; + + let mut events = Vec::new(); + + if let Some(text) = response.get_text() { + if !text.is_empty() { + events.push(StreamEvent::Delta(text)); + } + } + + if let Some(gf) = response.get_function() { + events.push(StreamEvent::FunctionCall(FunctionCall { + name: gf.name, + args: gf.args, + })); + } + + if let Some(usage) = &response.usage_metadata { + if let (Some(pt), Some(ct), Some(tt)) = ( + usage.prompt_token_count, + usage.candidates_token_count, + usage.total_token_count, + ) { + events.push(StreamEvent::Usage(Usage { + prompt_tokens: pt, + completion_tokens: ct, + total_tokens: tt, + })); + } + } + + events + } +} diff --git a/src/gemini/base.rs b/src/gemini/base.rs deleted file mode 100644 index 7a3a566..0000000 --- a/src/gemini/base.rs +++ /dev/null @@ -1,224 +0,0 @@ -use eventsource_stream::Eventsource; -use futures::{StreamExt, TryFutureExt, stream}; -use std::error::Error; - -use reqwest::RequestBuilder; - -use crate::{ - client::{ - Completion, FunctionCall, MessageType, Model, ModelRequest, Role, StreamEvent, - StreamResult, Usage, - }, - gemini::types::{ - Content, FunctionCallPart, FunctionResponsePart, GeminiRequest, GeminiResponse, GeminiTool, - GeminiTools, GenerationConfig, Part, SystemInstructionContent, ThinkingConfig, - }, -}; - -pub trait GeminiClient: Model { - fn create_request_body(&self, request: ModelRequest) -> GeminiRequest { - let thinking_config = request - .settings - .as_ref() - .and_then(|s| s.thinking_budget) - .map(|thinking_budget| ThinkingConfig { thinking_budget }); - - let generation_config = GenerationConfig { - max_output_tokens: request.settings.clone().and_then(|s| s.max_tokens), - temperature: request.settings.clone().and_then(|s| s.temperature), - thinking_config, - }; - - let contents: Vec = request - .messages - .clone() - .unwrap_or(vec![]) - .iter() - .map(|message| match &message.message_type { - MessageType::Text => Content { - parts: vec![Part::Text { - text: message.content.clone(), - }], - role: message.role.clone().unwrap_or_else(|| Role::User), - }, - MessageType::FunctionCall(fc) => Content { - parts: vec![Part::FunctionCall { - function_call: FunctionCallPart { - name: fc.name.clone(), - args: fc.args.clone(), - }, - }], - role: Role::Model, - }, - MessageType::FunctionResponse { name, response } => Content { - parts: vec![Part::FunctionResponse { - function_response: FunctionResponsePart { - name: name.clone(), - response: response.clone().unwrap_or(serde_json::Value::Null), - }, - }], - role: Role::User, - }, - }) - .collect(); - - let system_instruction = request.system.clone().map(|m| SystemInstructionContent { - parts: vec![Part::Text { text: m }], - }); - - let req = GeminiRequest { - system_instruction, - contents, - generation_config, - tools: request.tools.clone().map(|ts| { - vec![GeminiTools { - function_declarations: ts - .clone() - .iter() - .map(|t| GeminiTool::from_tool(t)) - .collect(), - }] - }), - }; - req - } - - async fn generate_content( - &self, - request: ModelRequest, - ) -> Result> { - let endpoint = self.get_endpoint(&self.model_name(), String::from("generateContent")); - let request_body = self.create_request_body(request); - let response = self - .build_request(&endpoint, &request_body) - .await? - .send() - .await?; - - let status = response.status(); - if !status.is_success() { - let error_text = response.text().map_err(|e| e.to_string()).await?; - return Err(format!( - "Gemini request failed with status {}: {}", - status, error_text - ) - .into()); - } - - let response_body: GeminiResponse = response.json().await?; - - let content: String = - response_body - .get_text() - .ok_or_else(|| -> Box { - "Missing completion from response".into() - })?; - - let prompt_tokens = response_body.get_prompt_tokens().unwrap_or(0); - let completion_tokens = response_body.get_completion_tokens().unwrap_or(0); - let total_tokens = response_body.get_total_tokens().unwrap_or(0); - - return Ok(Completion { - completion: content, - usage: Usage { - prompt_tokens, - completion_tokens, - total_tokens, - }, - function: response_body.get_function().map(|gf| FunctionCall { - name: gf.name, - args: gf.args, - }), - }); - } - - async fn stream_generate_content( - &self, - request: ModelRequest, - ) -> Result> { - let endpoint = self.get_endpoint( - &self.model_name(), - String::from("streamGenerateContent?alt=sse"), - ); - let request_body = self.create_request_body(request); - let response = self - .build_request(&endpoint, &request_body) - .await? - .send() - .await?; - - let status = response.status(); - if !status.is_success() { - let error_text = response.text().map_err(|e| e.to_string()).await?; - return Err(format!( - "Gemini streaming request failed with status {}: {}", - status, error_text - ) - .into()); - } - - let event_stream = response - .bytes_stream() - .eventsource() - .filter_map(|result| async { - match result { - Ok(event) => { - if event.data.is_empty() { - return None; - } - let parsed: Result = serde_json::from_str(&event.data); - match parsed { - Ok(gemini_response) => { - let mut events = Vec::new(); - - if let Some(text) = gemini_response.get_text() { - if !text.is_empty() { - events.push(StreamEvent::Delta(text)); - } - } - - if let Some(gf) = gemini_response.get_function() { - events.push(StreamEvent::FunctionCall(FunctionCall { - name: gf.name, - args: gf.args, - })); - } - - if let Some(usage) = &gemini_response.usage_metadata { - if let (Some(pt), Some(ct), Some(tt)) = ( - usage.prompt_token_count, - usage.candidates_token_count, - usage.total_token_count, - ) { - events.push(StreamEvent::Usage(Usage { - prompt_tokens: pt, - completion_tokens: ct, - total_tokens: tt, - })); - } - } - - if events.is_empty() { - None - } else { - Some(stream::iter(events)) - } - } - Err(e) => Some(stream::iter(vec![StreamEvent::Error(e.to_string())])), - } - } - Err(e) => Some(stream::iter(vec![StreamEvent::Error(e.to_string())])), - } - }) - .flat_map(|s| s); - - Ok(Box::pin(event_stream)) - } - - fn get_endpoint(&self, model: &String, method: String) -> String; - async fn build_request( - &self, - endpoint: &String, - request_body: &GeminiRequest, - ) -> Result>; -} diff --git a/src/gemini/direct_api_client.rs b/src/gemini/direct_api_client.rs deleted file mode 100644 index 6da5ba2..0000000 --- a/src/gemini/direct_api_client.rs +++ /dev/null @@ -1,61 +0,0 @@ -use std::error::Error; - -use crate::{ - client::{Completion, Model, ModelRequest, StreamResult}, - gemini::{ - base::GeminiClient, - types::{GeminiModel, GeminiRequest}, - }, -}; -use async_trait::async_trait; -use reqwest::RequestBuilder; - -pub struct GeminiApiModel { - pub api_key: String, - pub client: reqwest::Client, - pub model: GeminiModel, // TODO Replace this with a type -} - -#[async_trait] -impl Model for GeminiApiModel { - async fn completion( - &self, - request: ModelRequest, - ) -> Result> { - let response = self.generate_content(request).await?; - return Ok(response); - } - - async fn stream_completion( - &self, - request: ModelRequest, - ) -> Result> { - self.stream_generate_content(request).await - } - - fn model_name(&self) -> String { - self.model.to_string() - } -} - -impl GeminiClient for GeminiApiModel { - fn get_endpoint(&self, model: &String, method: String) -> String { - return format!( - "https://generativelanguage.googleapis.com/v1beta/models/{}:{}", - model, method - ); - } - - async fn build_request( - &self, - endpoint: &String, - request_body: &GeminiRequest, - ) -> Result> { - return Ok(self - .client - .post(endpoint.clone()) - .header("x-goog-api-key", self.api_key.clone()) - .header("Content-Type", "application/json") - .json(request_body)); - } -} diff --git a/src/gemini/mod.rs b/src/gemini/mod.rs index 7be025e..00b7820 100644 --- a/src/gemini/mod.rs +++ b/src/gemini/mod.rs @@ -1,12 +1,125 @@ -mod base; -mod direct_api_client; +mod adapter; mod gcloud_helpers; mod types; -mod vertex_client; #[cfg(test)] mod tests; -pub use direct_api_client::GeminiApiModel; +use async_trait::async_trait; + +use crate::provider::{Action, BoxError, LlmClient, Transport}; +use gcloud_helpers::get_access_token; + +pub use adapter::GeminiAdapter; pub use types::GeminiModel; -pub use vertex_client::GeminiVertexModel; + +pub type GeminiApiModel = LlmClient; + +pub type GeminiVertexModel = LlmClient; + +impl GeminiApiModel { + pub fn new(api_key: impl Into, model: GeminiModel) -> Self { + Self::with_client(api_key, model, reqwest::Client::new()) + } + + pub fn with_client( + api_key: impl Into, + model: GeminiModel, + client: reqwest::Client, + ) -> Self { + LlmClient::from_parts( + GeminiApiTransport { + api_key: api_key.into(), + client, + }, + model.to_string(), + ) + } +} + +impl GeminiVertexModel { + pub fn new(project_name: impl Into, model: GeminiModel) -> Self { + Self::with_client(project_name, model, reqwest::Client::new()) + } + + pub fn with_client( + project_name: impl Into, + model: GeminiModel, + client: reqwest::Client, + ) -> Self { + LlmClient::from_parts( + VertexTransport { + project_name: project_name.into(), + client, + }, + model.to_string(), + ) + } +} + +fn method(action: Action) -> &'static str { + match action { + Action::Generate => "generateContent", + Action::Stream => "streamGenerateContent?alt=sse", + } +} + +pub struct GeminiApiTransport { + pub api_key: String, + pub client: reqwest::Client, +} + +#[async_trait] +impl Transport for GeminiApiTransport { + async fn send( + &self, + model: &str, + action: Action, + body: serde_json::Value, + ) -> Result { + let url = format!( + "https://generativelanguage.googleapis.com/v1beta/models/{}:{}", + model, + method(action) + ); + Ok(self + .client + .post(url) + .header("x-goog-api-key", self.api_key.clone()) + .header("Content-Type", "application/json") + .json(&body) + .send() + .await?) + } +} + +pub struct VertexTransport { + pub project_name: String, + pub client: reqwest::Client, +} + +#[async_trait] +impl Transport for VertexTransport { + async fn send( + &self, + model: &str, + action: Action, + body: serde_json::Value, + ) -> Result { + let url = format!( + "https://aiplatform.googleapis.com/v1/projects/{}/locations/global/publishers/google/models/{}:{}", + self.project_name, + model, + method(action) + ); + let access_token = get_access_token().await?; + Ok(self + .client + .post(url) + .header("Authorization", format!("Bearer {}", access_token)) + .header("Content-Type", "application/json") + .json(&body) + .send() + .await?) + } +} diff --git a/src/gemini/tests.rs b/src/gemini/tests.rs index b0b052c..336ed0e 100644 --- a/src/gemini/tests.rs +++ b/src/gemini/tests.rs @@ -7,27 +7,25 @@ use serde::{Deserialize, Serialize}; use crate::{ client::{Message, Model, Settings, StreamEvent, Tool, Usage}, gemini::{ - base::GeminiClient, - direct_api_client::GeminiApiModel, + GeminiApiModel, GeminiVertexModel, + adapter::GeminiAdapter, types::{GeminiModel, GeminiTool}, - vertex_client::GeminiVertexModel, }, + provider::ProviderAdapter, }; fn make_direct(model: GeminiModel) -> GeminiApiModel { - GeminiApiModel { - client: reqwest::Client::new(), - api_key: env::var("GEMINI_KEY").expect("GEMINI_KEY env var must be set"), + GeminiApiModel::new( + env::var("GEMINI_KEY").expect("GEMINI_KEY env var must be set"), model, - } + ) } fn make_vertex(model: GeminiModel) -> GeminiVertexModel { - GeminiVertexModel { - project_name: env::var("VERTEX_PROJECT").expect("VERTEX_PROJECT env var must be set"), - client: reqwest::Client::new(), + GeminiVertexModel::new( + env::var("VERTEX_PROJECT").expect("VERTEX_PROJECT env var must be set"), model, - } + ) } fn default_settings() -> Settings { @@ -431,11 +429,11 @@ fn response_deserializes_when_candidate_has_no_content() { } fn make_direct_dummy(model: GeminiModel) -> GeminiApiModel { - GeminiApiModel { - client: reqwest::Client::new(), - api_key: "dummy".to_string(), - model, - } + GeminiApiModel::new("dummy", model) +} + +fn build_body(request: crate::client::ModelRequest) -> crate::gemini::types::GeminiRequest { + GeminiAdapter.build_body(&request, "gemini-test", false) } fn request_with_thinking(thinking_budget: Option) -> crate::client::ModelRequest { @@ -454,8 +452,7 @@ fn request_with_thinking(thinking_budget: Option) -> crate::client::ModelRe #[test] fn thinking_config_omitted_when_budget_is_none() { - let m = make_direct_dummy(GeminiModel::Gemini31Pro); - let body = m.create_request_body(request_with_thinking(None)); + let body = build_body(request_with_thinking(None)); assert!( body.generation_config.thinking_config.is_none(), "thinking_config should be omitted when thinking_budget is None" @@ -473,21 +470,19 @@ fn thinking_config_omitted_when_budget_is_none() { #[test] fn thinking_config_omitted_when_settings_is_none() { - let m = make_direct_dummy(GeminiModel::Gemini31Pro); let req = crate::client::ModelRequest { system: None, messages: Some(vec![Message::user("hi".to_string())]), settings: None, tools: None, }; - let body = m.create_request_body(req); + let body = build_body(req); assert!(body.generation_config.thinking_config.is_none()); } #[test] fn thinking_config_set_when_budget_is_some() { - let m = make_direct_dummy(GeminiModel::Gemini31Pro); - let body = m.create_request_body(request_with_thinking(Some(1024))); + let body = build_body(request_with_thinking(Some(1024))); let tc = body .generation_config .thinking_config @@ -508,8 +503,7 @@ fn thinking_config_set_when_budget_is_some() { #[test] fn thinking_config_supports_dynamic_budget() { // Gemini uses -1 to signal "dynamic thinking". Make sure we pass it through. - let m = make_direct_dummy(GeminiModel::Gemini31Pro); - let body = m.create_request_body(request_with_thinking(Some(-1))); + let body = build_body(request_with_thinking(Some(-1))); assert_eq!( body.generation_config .thinking_config @@ -536,17 +530,9 @@ fn test_model_name_gemini_api() { #[test] fn test_model_name_gemini_vertex() { - let m = GeminiVertexModel { - client: reqwest::Client::new(), - project_name: "dummy-project".to_string(), - model: GeminiModel::Gemini25Flash, - }; + let m = GeminiVertexModel::new("dummy-project", GeminiModel::Gemini25Flash); assert_eq!(m.model_name(), "gemini-2.5-flash"); - let m = GeminiVertexModel { - client: reqwest::Client::new(), - project_name: "dummy-project".to_string(), - model: GeminiModel::Gemini31Pro, - }; + let m = GeminiVertexModel::new("dummy-project", GeminiModel::Gemini31Pro); assert_eq!(m.model_name(), "gemini-3.1-pro-preview"); } diff --git a/src/gemini/vertex_client.rs b/src/gemini/vertex_client.rs deleted file mode 100644 index 05dc6b0..0000000 --- a/src/gemini/vertex_client.rs +++ /dev/null @@ -1,63 +0,0 @@ -use std::error::Error; - -use crate::{ - client::{Completion, Model, ModelRequest, StreamResult}, - gemini::{ - base::GeminiClient, - gcloud_helpers::get_access_token, - types::{GeminiModel, GeminiRequest}, - }, -}; -use async_trait::async_trait; -use reqwest::RequestBuilder; - -pub struct GeminiVertexModel { - pub project_name: String, - pub client: reqwest::Client, - pub model: GeminiModel, -} - -#[async_trait] -impl Model for GeminiVertexModel { - async fn completion( - &self, - request: ModelRequest, - ) -> Result> { - let response = self.generate_content(request).await?; - Ok(response) - } - - async fn stream_completion( - &self, - request: ModelRequest, - ) -> Result> { - self.stream_generate_content(request).await - } - - fn model_name(&self) -> String { - self.model.to_string() - } -} - -impl GeminiClient for GeminiVertexModel { - fn get_endpoint(&self, model: &String, method: String) -> String { - return format!( - "https://aiplatform.googleapis.com/v1/projects/{}/locations/global/publishers/google/models/{model}:{method}", - self.project_name - ); - } - - async fn build_request( - &self, - endpoint: &String, - request_body: &GeminiRequest, - ) -> Result> { - let access_token = get_access_token().await?; - return Ok(self - .client - .post(endpoint) - .header("Authorization", format!("Bearer {}", access_token)) - .header("Content-Type", "application/json") - .json(request_body)); - } -} diff --git a/src/lib.rs b/src/lib.rs index 49c0a53..7cd6453 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -2,6 +2,7 @@ pub mod claude; pub mod client; pub mod gemini; pub mod openai; +pub mod provider; pub use claude::{ClaudeApiModel, ClaudeModel}; pub use client::{ @@ -9,3 +10,4 @@ pub use client::{ }; pub use gemini::{GeminiApiModel, GeminiModel, GeminiVertexModel}; pub use openai::{OpenAiApiModel, OpenAiModel}; +pub use provider::{Action, LlmClient, ProviderAdapter, Transport}; diff --git a/src/openai/base.rs b/src/openai/adapter.rs similarity index 51% rename from src/openai/base.rs rename to src/openai/adapter.rs index ef69768..aab8446 100644 --- a/src/openai/base.rs +++ b/src/openai/adapter.rs @@ -1,34 +1,27 @@ use std::collections::HashMap; -use std::error::Error; - -use eventsource_stream::Eventsource; -use futures::{StreamExt, TryFutureExt, stream}; -use reqwest::RequestBuilder; use crate::{ - client::{ - Completion, FunctionCall, MessageType, Model, ModelRequest, StreamEvent, StreamResult, - Usage, - }, + client::{Completion, FunctionCall, MessageType, ModelRequest, StreamEvent, Usage}, openai::types::{ OpenAiInputItem, OpenAiRequest, OpenAiResponse, OpenAiTool, ResponsesStreamEvent, synth_call_id, }, + provider::{BoxError, ProviderAdapter}, }; -pub trait OpenAiClient: Model { - fn create_request_body(&self, request: ModelRequest, stream: bool) -> OpenAiRequest { - let settings = request.settings.clone(); +#[derive(Debug, Default, Clone, Copy)] +pub struct OpenAiAdapter; - let max_output_tokens = settings - .as_ref() - .and_then(|s| s.max_tokens) - .map(|v| v as i32); +impl ProviderAdapter for OpenAiAdapter { + const NAME: &'static str = "OpenAI"; + type Request = OpenAiRequest; + type StreamState = (); - let temperature = settings - .as_ref() - .and_then(|s| s.temperature) - .map(|v| v as f32); + fn build_body(&self, request: &ModelRequest, model: &str, stream: bool) -> OpenAiRequest { + let settings = request.settings.as_ref(); + + let max_output_tokens = settings.and_then(|s| s.max_tokens).map(|v| v as i32); + let temperature = settings.and_then(|s| s.temperature); // Build input items (no system message — that goes to `instructions`). let mut input: Vec = Vec::new(); @@ -71,38 +64,23 @@ pub trait OpenAiClient: Model { let tools = request .tools - .clone() + .as_ref() .map(|ts| ts.iter().map(OpenAiTool::from_tool).collect()); - let stream_flag = if stream { Some(true) } else { None }; - OpenAiRequest { - model: self.model_name(), + model: model.to_string(), input, instructions: request.system.clone(), max_output_tokens, temperature, tools, - stream: stream_flag, + stream: if stream { Some(true) } else { None }, store: false, } } - async fn generate_content( - &self, - request: ModelRequest, - ) -> Result> { - let endpoint = self.get_endpoint(); - let body = self.create_request_body(request, false); - let response = self.build_request(&endpoint, &body).await?.send().await?; - - let status = response.status(); - if !status.is_success() { - let err = response.text().map_err(|e| e.to_string()).await?; - return Err(format!("OpenAI request failed with status {}: {}", status, err).into()); - } - - let body: OpenAiResponse = response.json().await?; + fn parse_completion(&self, body: &[u8]) -> Result { + let body: OpenAiResponse = serde_json::from_slice(body)?; let text = body.get_text(); let function = body @@ -126,74 +104,25 @@ pub trait OpenAiClient: Model { }) } - async fn stream_generate_content( - &self, - request: ModelRequest, - ) -> Result> { - let endpoint = self.get_endpoint(); - let body = self.create_request_body(request, true); - let response = self.build_request(&endpoint, &body).await?.send().await?; - - let status = response.status(); - if !status.is_success() { - let err = response.text().map_err(|e| e.to_string()).await?; - return Err(format!( - "OpenAI streaming request failed with status {}: {}", - status, err - ) - .into()); + fn map_sse_event(&self, data: &str, _state: &mut ()) -> Vec { + if data == "[DONE]" { + return Vec::new(); } - - let sse = Box::pin(response.bytes_stream().eventsource()); - let state = State { - sse, - buffer: std::collections::VecDeque::new(), - }; - - let out = stream::unfold(state, |mut state| async move { - loop { - if let Some(ev) = state.buffer.pop_front() { - return Some((ev, state)); - } - - let next = state.sse.next().await?; - match next { - Err(e) => { - state.buffer.push_back(StreamEvent::Error(e.to_string())); - } - Ok(event) => { - if event.data.is_empty() { - continue; - } - let parsed: Result = - serde_json::from_str(&event.data); - match parsed { - Err(e) => state.buffer.push_back(StreamEvent::Error(e.to_string())), - Ok(ev) => handle_stream_event(ev, &mut state), - } - } - } - } - }); - - Ok(Box::pin(out)) + let mut out = Vec::new(); + match serde_json::from_str::(data) { + Err(e) => out.push(StreamEvent::Error(e.to_string())), + Ok(ev) => handle_event(ev, &mut out), + } + out } - - fn get_endpoint(&self) -> String; - - async fn build_request( - &self, - endpoint: &String, - request_body: &OpenAiRequest, - ) -> Result>; } -fn handle_stream_event(event: ResponsesStreamEvent, state: &mut State) { +fn handle_event(event: ResponsesStreamEvent, out: &mut Vec) { match event.event_type.as_str() { "response.output_text.delta" => { if let Some(delta) = event.delta { if !delta.is_empty() { - state.push_event(StreamEvent::Delta(delta)); + out.push(StreamEvent::Delta(delta)); } } } @@ -208,7 +137,7 @@ fn handle_stream_event(event: ResponsesStreamEvent, state: &mut State) { match serde_json::from_str(&args_str) { Ok(v) => v, Err(e) => { - state.push_event(StreamEvent::Error(format!( + out.push(StreamEvent::Error(format!( "failed to parse streamed tool arguments JSON: {}", e ))); @@ -216,7 +145,7 @@ fn handle_stream_event(event: ResponsesStreamEvent, state: &mut State) { } } }; - state.push_event(StreamEvent::FunctionCall(FunctionCall { name, args })); + out.push(StreamEvent::FunctionCall(FunctionCall { name, args })); } } } @@ -224,7 +153,7 @@ fn handle_stream_event(event: ResponsesStreamEvent, state: &mut State) { "response.completed" => { if let Some(resp) = event.response { if let Some(usage) = resp.usage { - state.push_event(StreamEvent::Usage(Usage { + out.push(StreamEvent::Usage(Usage { prompt_tokens: usage.input_tokens, completion_tokens: usage.output_tokens, total_tokens: usage.total_tokens, @@ -237,23 +166,3 @@ fn handle_stream_event(event: ResponsesStreamEvent, state: &mut State) { } } } - -struct State { - sse: std::pin::Pin< - Box< - dyn futures::Stream< - Item = Result< - eventsource_stream::Event, - eventsource_stream::EventStreamError, - >, - > + Send, - >, - >, - buffer: std::collections::VecDeque, -} - -impl State { - fn push_event(&mut self, ev: StreamEvent) { - self.buffer.push_back(ev); - } -} diff --git a/src/openai/direct_api_client.rs b/src/openai/direct_api_client.rs deleted file mode 100644 index 5b8f3ea..0000000 --- a/src/openai/direct_api_client.rs +++ /dev/null @@ -1,58 +0,0 @@ -use std::error::Error; - -use async_trait::async_trait; -use reqwest::RequestBuilder; - -use crate::{ - client::{Completion, Model, ModelRequest, StreamResult}, - openai::{ - base::OpenAiClient, - types::{OpenAiModel, OpenAiRequest}, - }, -}; - -pub struct OpenAiApiModel { - pub api_key: String, - pub client: reqwest::Client, - pub model: OpenAiModel, -} - -#[async_trait] -impl Model for OpenAiApiModel { - async fn completion( - &self, - request: ModelRequest, - ) -> Result> { - self.generate_content(request).await - } - - async fn stream_completion( - &self, - request: ModelRequest, - ) -> Result> { - self.stream_generate_content(request).await - } - - fn model_name(&self) -> String { - self.model.to_string() - } -} - -impl OpenAiClient for OpenAiApiModel { - fn get_endpoint(&self) -> String { - "https://api.openai.com/v1/responses".to_string() - } - - async fn build_request( - &self, - endpoint: &String, - request_body: &OpenAiRequest, - ) -> Result> { - Ok(self - .client - .post(endpoint) - .header("Authorization", format!("Bearer {}", self.api_key)) - .header("Content-Type", "application/json") - .json(request_body)) - } -} diff --git a/src/openai/mod.rs b/src/openai/mod.rs index 10bca29..8063ad0 100644 --- a/src/openai/mod.rs +++ b/src/openai/mod.rs @@ -1,9 +1,58 @@ -mod base; -mod direct_api_client; +mod adapter; mod types; #[cfg(test)] mod tests; -pub use direct_api_client::OpenAiApiModel; +use async_trait::async_trait; + +use crate::provider::{Action, BoxError, LlmClient, Transport}; + +pub use adapter::OpenAiAdapter; pub use types::OpenAiModel; + +pub type OpenAiApiModel = LlmClient; + +impl OpenAiApiModel { + pub fn new(api_key: impl Into, model: OpenAiModel) -> Self { + Self::with_client(api_key, model, reqwest::Client::new()) + } + + pub fn with_client( + api_key: impl Into, + model: OpenAiModel, + client: reqwest::Client, + ) -> Self { + LlmClient::from_parts( + OpenAiTransport { + api_key: api_key.into(), + client, + }, + model.to_string(), + ) + } +} + +pub struct OpenAiTransport { + pub api_key: String, + pub client: reqwest::Client, +} + +#[async_trait] +impl Transport for OpenAiTransport { + async fn send( + &self, + _model: &str, + _action: Action, + body: serde_json::Value, + ) -> Result { + Ok(self + .client + .post("https://api.openai.com/v1/responses") + .header("Authorization", format!("Bearer {}", self.api_key)) + .header("Content-Type", "application/json") + .json(&body) + .send() + .await?) + } +} diff --git a/src/openai/tests.rs b/src/openai/tests.rs index d73cd3e..7bc971c 100644 --- a/src/openai/tests.rs +++ b/src/openai/tests.rs @@ -7,17 +7,16 @@ use serde::{Deserialize, Serialize}; use crate::{ client::{Message, Model, Settings, StreamEvent, Tool, Usage}, openai::{ - direct_api_client::OpenAiApiModel, + OpenAiApiModel, types::{OpenAiModel, OpenAiTool}, }, }; fn make_model(model: OpenAiModel) -> OpenAiApiModel { - OpenAiApiModel { - client: reqwest::Client::new(), - api_key: env::var("OPENAI_KEY").expect("OPENAI_KEY env var must be set"), + OpenAiApiModel::new( + env::var("OPENAI_KEY").expect("OPENAI_KEY env var must be set"), model, - } + ) } fn default_settings() -> Settings { @@ -346,38 +345,18 @@ fn test_openai_tool_emits_required_and_name_and_description() { #[test] fn test_model_name_openai_api() { - let m = OpenAiApiModel { - client: reqwest::Client::new(), - api_key: "dummy-key".to_string(), - model: OpenAiModel::Gpt5_4, - }; + let m = OpenAiApiModel::new("dummy-key", OpenAiModel::Gpt5_4); assert_eq!(m.model_name(), "gpt-5.4"); - let m = OpenAiApiModel { - client: reqwest::Client::new(), - api_key: "dummy-key".to_string(), - model: OpenAiModel::Gpt5_4Mini, - }; + let m = OpenAiApiModel::new("dummy-key", OpenAiModel::Gpt5_4Mini); assert_eq!(m.model_name(), "gpt-5.4-mini"); - let m = OpenAiApiModel { - client: reqwest::Client::new(), - api_key: "dummy-key".to_string(), - model: OpenAiModel::Gpt5_4Nano, - }; + let m = OpenAiApiModel::new("dummy-key", OpenAiModel::Gpt5_4Nano); assert_eq!(m.model_name(), "gpt-5.4-nano"); - let m = OpenAiApiModel { - client: reqwest::Client::new(), - api_key: "dummy-key".to_string(), - model: OpenAiModel::Gpt5_5, - }; + let m = OpenAiApiModel::new("dummy-key", OpenAiModel::Gpt5_5); assert_eq!(m.model_name(), "gpt-5.5"); - let m = OpenAiApiModel { - client: reqwest::Client::new(), - api_key: "dummy-key".to_string(), - model: OpenAiModel::Gpt5_3Codex, - }; + let m = OpenAiApiModel::new("dummy-key", OpenAiModel::Gpt5_3Codex); assert_eq!(m.model_name(), "gpt-5.3-codex"); } diff --git a/src/provider.rs b/src/provider.rs new file mode 100644 index 0000000..c60d91a --- /dev/null +++ b/src/provider.rs @@ -0,0 +1,170 @@ +//! Provider abstraction: one generic client composed from two small traits. +//! +//! - [`ProviderAdapter`] — pure request/response conversions for one provider +//! wire format (no I/O). Trivially unit-testable. +//! - [`Transport`] — endpoint + authentication for one provider backend +//! (only I/O). +//! - [`LlmClient`] — composes the two and implements [`Model`] exactly once. +//! +//! Concrete client types like `ClaudeApiModel` are type aliases of +//! [`LlmClient`] with the right adapter/transport pair. + +use std::collections::VecDeque; +use std::error::Error; +use std::pin::Pin; + +use async_trait::async_trait; +use eventsource_stream::Eventsource; +use futures::{StreamExt, stream}; +use serde::Serialize; + +use crate::client::{Completion, Model, ModelRequest, StreamEvent, StreamResult}; + +pub type BoxError = Box; + +/// Which provider action a request is for. Transports use this to select the +/// endpoint when the URL differs between plain and streaming completions +/// (e.g. Gemini's `generateContent` vs `streamGenerateContent`). +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum Action { + Generate, + Stream, +} + +/// Pure request/response conversions for one provider wire format. +/// +/// Implementations are stateless unit structs; per-stream accumulation +/// (e.g. partially streamed tool-call JSON) lives in [`Self::StreamState`]. +pub trait ProviderAdapter: Default + Send + Sync + 'static { + /// Provider name used in error messages. + const NAME: &'static str; + /// Wire request body type. + type Request: Serialize; + /// State threaded through an SSE stream (use `()` when none is needed). + type StreamState: Default + Send + 'static; + + /// Convert the common request into the provider's wire request. + fn build_body(&self, request: &ModelRequest, model: &str, stream: bool) -> Self::Request; + + /// Parse a successful non-streaming response body. + fn parse_completion(&self, body: &[u8]) -> Result; + + /// Map one SSE `data:` payload to zero or more common stream events. + fn map_sse_event(&self, data: &str, state: &mut Self::StreamState) -> Vec; +} + +/// Endpoint + authentication for one provider backend. +#[async_trait] +pub trait Transport: Send + Sync { + /// Send `body` as an authenticated JSON POST for `model`/`action` and + /// return the raw response (status is checked by the caller). + async fn send( + &self, + model: &str, + action: Action, + body: serde_json::Value, + ) -> Result; +} + +/// Generic LLM client: a [`ProviderAdapter`] (wire format) plus a +/// [`Transport`] (endpoint + auth). Implements [`Model`] once for all +/// providers. +pub struct LlmClient { + pub adapter: A, + pub transport: T, + pub model: String, +} + +impl LlmClient { + pub fn from_parts(transport: T, model: impl Into) -> Self { + LlmClient { + adapter: A::default(), + transport, + model: model.into(), + } + } +} + +#[async_trait] +impl Model for LlmClient { + async fn completion(&self, request: ModelRequest) -> Result { + let body = serde_json::to_value(self.adapter.build_body(&request, &self.model, false))?; + let response = self.transport.send(&self.model, Action::Generate, body).await?; + let response = check_status(A::NAME, response).await?; + let bytes = response.bytes().await?; + self.adapter.parse_completion(&bytes) + } + + async fn stream_completion(&self, request: ModelRequest) -> Result { + let body = serde_json::to_value(self.adapter.build_body(&request, &self.model, true))?; + let response = self.transport.send(&self.model, Action::Stream, body).await?; + let response = check_status(A::NAME, response).await?; + Ok(sse_stream(A::default(), response)) + } + + fn model_name(&self) -> String { + self.model.clone() + } +} + +/// Pass through successful responses; turn error statuses into a readable error. +async fn check_status( + provider: &str, + response: reqwest::Response, +) -> Result { + let status = response.status(); + if status.is_success() { + return Ok(response); + } + let body = response.text().await.unwrap_or_default(); + Err(format!("{} request failed with status {}: {}", provider, status, body).into()) +} + +type SseEvents = Pin< + Box< + dyn futures::Stream< + Item = Result< + eventsource_stream::Event, + eventsource_stream::EventStreamError, + >, + > + Send, + >, +>; + +/// Turn an SSE response into a `StreamResult` by feeding each `data:` payload +/// through the adapter's [`ProviderAdapter::map_sse_event`]. +fn sse_stream(adapter: A, response: reqwest::Response) -> StreamResult { + struct State { + adapter: A, + provider_state: A::StreamState, + buffer: VecDeque, + sse: SseEvents, + } + + let state = State { + adapter, + provider_state: A::StreamState::default(), + buffer: VecDeque::new(), + sse: Box::pin(response.bytes_stream().eventsource()), + }; + + let out = stream::unfold(state, |mut st| async move { + loop { + if let Some(ev) = st.buffer.pop_front() { + return Some((ev, st)); + } + match st.sse.next().await? { + Err(e) => st.buffer.push_back(StreamEvent::Error(e.to_string())), + Ok(event) => { + if event.data.is_empty() { + continue; + } + let events = st.adapter.map_sse_event(&event.data, &mut st.provider_state); + st.buffer.extend(events); + } + } + } + }); + + Box::pin(out) +} From bae6b1448e896eb20cb5d73623a994934dd7d00b Mon Sep 17 00:00:00 2001 From: Endi Sukaj Date: Sun, 30 Aug 2026 23:02:57 +0200 Subject: [PATCH 3/3] remove duplicated code --- README.md | 17 ++ src/error.rs | 74 +++++++++ src/http.rs | 84 ++++++++++ src/lib.rs | 27 ++-- src/provider.rs | 95 ++--------- src/{ => providers}/claude/adapter.rs | 21 +-- src/{ => providers}/claude/mod.rs | 5 +- src/{ => providers}/claude/tests.rs | 5 +- src/{ => providers}/claude/types.rs | 30 +--- src/{ => providers}/gemini/adapter.rs | 22 ++- src/{ => providers}/gemini/gcloud_helpers.rs | 0 src/{ => providers}/gemini/mod.rs | 9 +- src/{ => providers}/gemini/tests.rs | 19 +-- src/{ => providers}/gemini/types.rs | 2 +- src/providers/mod.rs | 3 + src/{ => providers}/openai/adapter.rs | 19 ++- src/{ => providers}/openai/mod.rs | 5 +- src/{ => providers}/openai/tests.rs | 5 +- src/{ => providers}/openai/types.rs | 27 +--- src/request.rs | 97 +++++++++++ src/{client => }/tests.rs | 6 +- src/{client/mod.rs => types.rs} | 162 ++++--------------- 22 files changed, 419 insertions(+), 315 deletions(-) create mode 100644 src/error.rs create mode 100644 src/http.rs rename src/{ => providers}/claude/adapter.rs (93%) rename src/{ => providers}/claude/mod.rs (91%) rename src/{ => providers}/claude/tests.rs (99%) rename src/{ => providers}/claude/types.rs (82%) rename src/{ => providers}/gemini/adapter.rs (87%) rename src/{ => providers}/gemini/gcloud_helpers.rs (100%) rename src/{ => providers}/gemini/mod.rs (92%) rename src/{ => providers}/gemini/tests.rs (97%) rename src/{ => providers}/gemini/types.rs (99%) create mode 100644 src/providers/mod.rs rename src/{ => providers}/openai/adapter.rs (94%) rename src/{ => providers}/openai/mod.rs (90%) rename src/{ => providers}/openai/tests.rs (99%) rename src/{ => providers}/openai/types.rs (85%) create mode 100644 src/request.rs rename src/{client => }/tests.rs (98%) rename src/{client/mod.rs => types.rs} (52%) diff --git a/README.md b/README.md index 70a717b..18fe64d 100644 --- a/README.md +++ b/README.md @@ -285,6 +285,21 @@ To use a custom `reqwest::Client`, call the `with_client` constructors; to add a new backend, implement `Transport` (new auth/endpoint for an existing wire format) or `ProviderAdapter` + `Transport` (entirely new provider). +Module layout: + +``` +src/ + error.rs // LlmError + types.rs // Message, Role, Tool, Settings, Usage, Completion, StreamEvent + request.rs // Model trait, ModelRequest + builder + http.rs // shared status-check + SSE-to-StreamEvent plumbing + provider.rs // ProviderAdapter, Transport, LlmClient + providers/ + claude/ // wire types + adapter + transport + openai/ + gemini/ // + Vertex transport, gcloud helpers +``` + ## Core types cheat-sheet - `Model` — trait with `completion()` and `stream_completion()`; all providers implement it. @@ -296,6 +311,8 @@ format) or `ProviderAdapter` + `Transport` (entirely new provider). - `Settings { max_tokens, timeout, temperature, thinking_budget }` — all `Option`. - `Completion { completion, usage, function }` — unified non-streaming response. - `StreamEvent` — `Delta | Usage | FunctionCall | Error` for streaming. +- `LlmError` — typed error for all providers: `Http { provider, status, body } | + Parse | Transport | Auth | Other`. Converts into `Box` via `?`. ## Known limitations diff --git a/src/error.rs b/src/error.rs new file mode 100644 index 0000000..31ab392 --- /dev/null +++ b/src/error.rs @@ -0,0 +1,74 @@ +use std::error::Error; +use std::fmt; + +/// Unified error type for all providers. +/// +/// Implemented by hand instead of via `thiserror` to keep the dependency +/// tree minimal; the shape is the same. +#[derive(Debug)] +pub enum LlmError { + /// The provider returned a non-success HTTP status. + Http { + provider: &'static str, + status: u16, + body: String, + }, + /// A request or response body could not be (de)serialized. + Parse(serde_json::Error), + /// The HTTP request itself failed (connection, TLS, timeout, ...). + Transport(reqwest::Error), + /// Authentication failed (e.g. obtaining a gcloud access token). + Auth(String), + /// Anything else (e.g. a well-formed response missing required fields). + Other(String), +} + +impl fmt::Display for LlmError { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + match self { + LlmError::Http { + provider, + status, + body, + } => write!(f, "{} request failed with status {}: {}", provider, status, body), + LlmError::Parse(e) => write!(f, "failed to parse provider payload: {}", e), + LlmError::Transport(e) => write!(f, "transport error: {}", e), + LlmError::Auth(msg) => write!(f, "authentication error: {}", msg), + LlmError::Other(msg) => write!(f, "{}", msg), + } + } +} + +impl Error for LlmError { + fn source(&self) -> Option<&(dyn Error + 'static)> { + match self { + LlmError::Parse(e) => Some(e), + LlmError::Transport(e) => Some(e), + _ => None, + } + } +} + +impl From for LlmError { + fn from(e: serde_json::Error) -> Self { + LlmError::Parse(e) + } +} + +impl From for LlmError { + fn from(e: reqwest::Error) -> Self { + LlmError::Transport(e) + } +} + +impl From for LlmError { + fn from(msg: String) -> Self { + LlmError::Other(msg) + } +} + +impl From<&str> for LlmError { + fn from(msg: &str) -> Self { + LlmError::Other(msg.to_string()) + } +} diff --git a/src/http.rs b/src/http.rs new file mode 100644 index 0000000..860491b --- /dev/null +++ b/src/http.rs @@ -0,0 +1,84 @@ +//! Shared HTTP/SSE plumbing used by every provider: status checking and the +//! SSE-to-`StreamEvent` pump. Exists exactly once. + +use std::collections::VecDeque; +use std::pin::Pin; + +use eventsource_stream::Eventsource; +use futures::{StreamExt, stream}; + +use crate::error::LlmError; +use crate::types::{StreamEvent, StreamResult}; + +/// Pass successful responses through; turn error statuses into [`LlmError::Http`]. +pub async fn require_success( + provider: &'static str, + response: reqwest::Response, +) -> Result { + let status = response.status(); + if status.is_success() { + return Ok(response); + } + let body = response.text().await.unwrap_or_default(); + Err(LlmError::Http { + provider, + status: status.as_u16(), + body, + }) +} + +type SseSource = Pin< + Box< + dyn futures::Stream< + Item = Result< + eventsource_stream::Event, + eventsource_stream::EventStreamError, + >, + > + Send, + >, +>; + +/// Turn an SSE response into a `StreamResult`. +/// +/// Each non-empty `data:` payload is passed to `handler` together with the +/// caller's accumulation state; the handler returns zero or more common +/// stream events. Transport errors surface as [`StreamEvent::Error`]. +pub fn sse_events( + response: reqwest::Response, + state: S, + handler: impl Fn(&str, &mut S) -> Vec + Send + 'static, +) -> StreamResult { + struct Pump { + sse: SseSource, + state: S, + handler: F, + buffer: VecDeque, + } + + let pump = Pump { + sse: Box::pin(response.bytes_stream().eventsource()), + state, + handler, + buffer: VecDeque::new(), + }; + + let out = stream::unfold(pump, |mut p| async move { + loop { + if let Some(ev) = p.buffer.pop_front() { + return Some((ev, p)); + } + match p.sse.next().await? { + Err(e) => p.buffer.push_back(StreamEvent::Error(e.to_string())), + Ok(event) => { + if event.data.is_empty() { + continue; + } + let events = (p.handler)(&event.data, &mut p.state); + p.buffer.extend(events); + } + } + } + }); + + Box::pin(out) +} diff --git a/src/lib.rs b/src/lib.rs index 7cd6453..38e29f1 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -1,13 +1,20 @@ -pub mod claude; -pub mod client; -pub mod gemini; -pub mod openai; +pub mod error; +pub mod http; pub mod provider; +pub mod providers; +pub mod request; +pub mod types; -pub use claude::{ClaudeApiModel, ClaudeModel}; -pub use client::{ - Message, MessageType, ModelRequest, Role, Settings, StreamEvent, StreamResult, Tool, -}; -pub use gemini::{GeminiApiModel, GeminiModel, GeminiVertexModel}; -pub use openai::{OpenAiApiModel, OpenAiModel}; +#[cfg(test)] +mod tests; + +pub use error::LlmError; pub use provider::{Action, LlmClient, ProviderAdapter, Transport}; +pub use providers::claude::{ClaudeApiModel, ClaudeModel}; +pub use providers::gemini::{GeminiApiModel, GeminiModel, GeminiVertexModel}; +pub use providers::openai::{OpenAiApiModel, OpenAiModel}; +pub use request::{Model, ModelRequest, ModelRequestBuilder}; +pub use types::{ + Completion, FunctionCall, Message, MessageType, Role, Settings, StreamEvent, StreamResult, + Tool, ToolParameters, Usage, +}; diff --git a/src/provider.rs b/src/provider.rs index c60d91a..53ae2b7 100644 --- a/src/provider.rs +++ b/src/provider.rs @@ -9,18 +9,13 @@ //! Concrete client types like `ClaudeApiModel` are type aliases of //! [`LlmClient`] with the right adapter/transport pair. -use std::collections::VecDeque; -use std::error::Error; -use std::pin::Pin; - use async_trait::async_trait; -use eventsource_stream::Eventsource; -use futures::{StreamExt, stream}; use serde::Serialize; -use crate::client::{Completion, Model, ModelRequest, StreamEvent, StreamResult}; - -pub type BoxError = Box; +use crate::error::LlmError; +use crate::http::{require_success, sse_events}; +use crate::request::{Model, ModelRequest}; +use crate::types::{Completion, StreamEvent, StreamResult}; /// Which provider action a request is for. Transports use this to select the /// endpoint when the URL differs between plain and streaming completions @@ -47,7 +42,7 @@ pub trait ProviderAdapter: Default + Send + Sync + 'static { fn build_body(&self, request: &ModelRequest, model: &str, stream: bool) -> Self::Request; /// Parse a successful non-streaming response body. - fn parse_completion(&self, body: &[u8]) -> Result; + fn parse_completion(&self, body: &[u8]) -> Result; /// Map one SSE `data:` payload to zero or more common stream events. fn map_sse_event(&self, data: &str, state: &mut Self::StreamState) -> Vec; @@ -63,7 +58,7 @@ pub trait Transport: Send + Sync { model: &str, action: Action, body: serde_json::Value, - ) -> Result; + ) -> Result; } /// Generic LLM client: a [`ProviderAdapter`] (wire format) plus a @@ -87,84 +82,28 @@ impl LlmClient { #[async_trait] impl Model for LlmClient { - async fn completion(&self, request: ModelRequest) -> Result { + async fn completion(&self, request: ModelRequest) -> Result { let body = serde_json::to_value(self.adapter.build_body(&request, &self.model, false))?; let response = self.transport.send(&self.model, Action::Generate, body).await?; - let response = check_status(A::NAME, response).await?; + let response = require_success(A::NAME, response).await?; let bytes = response.bytes().await?; self.adapter.parse_completion(&bytes) } - async fn stream_completion(&self, request: ModelRequest) -> Result { + async fn stream_completion(&self, request: ModelRequest) -> Result { let body = serde_json::to_value(self.adapter.build_body(&request, &self.model, true))?; let response = self.transport.send(&self.model, Action::Stream, body).await?; - let response = check_status(A::NAME, response).await?; - Ok(sse_stream(A::default(), response)) + let response = require_success(A::NAME, response).await?; + + let adapter = A::default(); + Ok(sse_events( + response, + A::StreamState::default(), + move |data, state| adapter.map_sse_event(data, state), + )) } fn model_name(&self) -> String { self.model.clone() } } - -/// Pass through successful responses; turn error statuses into a readable error. -async fn check_status( - provider: &str, - response: reqwest::Response, -) -> Result { - let status = response.status(); - if status.is_success() { - return Ok(response); - } - let body = response.text().await.unwrap_or_default(); - Err(format!("{} request failed with status {}: {}", provider, status, body).into()) -} - -type SseEvents = Pin< - Box< - dyn futures::Stream< - Item = Result< - eventsource_stream::Event, - eventsource_stream::EventStreamError, - >, - > + Send, - >, ->; - -/// Turn an SSE response into a `StreamResult` by feeding each `data:` payload -/// through the adapter's [`ProviderAdapter::map_sse_event`]. -fn sse_stream(adapter: A, response: reqwest::Response) -> StreamResult { - struct State { - adapter: A, - provider_state: A::StreamState, - buffer: VecDeque, - sse: SseEvents, - } - - let state = State { - adapter, - provider_state: A::StreamState::default(), - buffer: VecDeque::new(), - sse: Box::pin(response.bytes_stream().eventsource()), - }; - - let out = stream::unfold(state, |mut st| async move { - loop { - if let Some(ev) = st.buffer.pop_front() { - return Some((ev, st)); - } - match st.sse.next().await? { - Err(e) => st.buffer.push_back(StreamEvent::Error(e.to_string())), - Ok(event) => { - if event.data.is_empty() { - continue; - } - let events = st.adapter.map_sse_event(&event.data, &mut st.provider_state); - st.buffer.extend(events); - } - } - } - }); - - Box::pin(out) -} diff --git a/src/claude/adapter.rs b/src/providers/claude/adapter.rs similarity index 93% rename from src/claude/adapter.rs rename to src/providers/claude/adapter.rs index d10fd2d..ee69b76 100644 --- a/src/claude/adapter.rs +++ b/src/providers/claude/adapter.rs @@ -1,13 +1,16 @@ use std::collections::HashMap; use crate::{ - claude::types::{ - BlockDelta, ClaudeMessage, ClaudeRequest, ClaudeResponse, ClaudeTool, ContentBlock, - DEFAULT_MAX_TOKENS, ResponseBlock, StreamContentBlock, StreamingEvent, ThinkingConfig, - synth_tool_use_id, - }, - client::{Completion, FunctionCall, MessageType, ModelRequest, StreamEvent, Usage}, - provider::{BoxError, ProviderAdapter}, + error::LlmError, + provider::ProviderAdapter, + request::ModelRequest, + types::{Completion, FunctionCall, MessageType, StreamEvent, Usage}, +}; + +use super::types::{ + BlockDelta, ClaudeMessage, ClaudeRequest, ClaudeResponse, ClaudeTool, ContentBlock, + DEFAULT_MAX_TOKENS, ResponseBlock, StreamContentBlock, StreamingEvent, ThinkingConfig, + synth_tool_use_id, }; /// Pure conversions between the common request/response types and the @@ -46,7 +49,7 @@ impl ProviderAdapter for ClaudeAdapter { .map(|m| match &m.message_type { MessageType::Text => ClaudeMessage { role: match m.role { - Some(crate::client::Role::Model) => "assistant", + Some(crate::types::Role::Model) => "assistant", _ => "user", }, content: vec![ContentBlock::Text { @@ -91,7 +94,7 @@ impl ProviderAdapter for ClaudeAdapter { } } - fn parse_completion(&self, body: &[u8]) -> Result { + fn parse_completion(&self, body: &[u8]) -> Result { let body: ClaudeResponse = serde_json::from_slice(body)?; let mut text = String::new(); diff --git a/src/claude/mod.rs b/src/providers/claude/mod.rs similarity index 91% rename from src/claude/mod.rs rename to src/providers/claude/mod.rs index 857af07..04a3f08 100644 --- a/src/claude/mod.rs +++ b/src/providers/claude/mod.rs @@ -6,7 +6,8 @@ mod tests; use async_trait::async_trait; -use crate::provider::{Action, BoxError, LlmClient, Transport}; +use crate::error::LlmError; +use crate::provider::{Action, LlmClient, Transport}; pub use adapter::ClaudeAdapter; pub use types::ClaudeModel; @@ -46,7 +47,7 @@ impl Transport for AnthropicTransport { _model: &str, _action: Action, body: serde_json::Value, - ) -> Result { + ) -> Result { Ok(self .client .post("https://api.anthropic.com/v1/messages") diff --git a/src/claude/tests.rs b/src/providers/claude/tests.rs similarity index 99% rename from src/claude/tests.rs rename to src/providers/claude/tests.rs index fa526ac..bae52b9 100644 --- a/src/claude/tests.rs +++ b/src/providers/claude/tests.rs @@ -5,11 +5,12 @@ use schemars::JsonSchema; use serde::{Deserialize, Serialize}; use crate::{ - claude::{ + providers::claude::{ ClaudeApiModel, types::{ClaudeModel, ClaudeTool}, }, - client::{Message, Model, Settings, StreamEvent, Tool, Usage}, + request::Model, + types::{Message, Settings, StreamEvent, Tool, Usage}, }; fn make_model(model: ClaudeModel) -> ClaudeApiModel { diff --git a/src/claude/types.rs b/src/providers/claude/types.rs similarity index 82% rename from src/claude/types.rs rename to src/providers/claude/types.rs index 03718bb..6dde4d5 100644 --- a/src/claude/types.rs +++ b/src/providers/claude/types.rs @@ -1,6 +1,6 @@ use std::collections::HashMap; -use crate::client::Tool; +use crate::types::Tool; use serde::{Deserialize, Serialize}; use serde_json::Value; @@ -75,35 +75,11 @@ pub struct ClaudeTool { impl ClaudeTool { pub fn from_tool(tool: &Tool) -> ClaudeTool { - // Anthropic accepts standard JSON Schema. If the caller provided - // `ToolParameters`, serialize them straight through; otherwise emit - // an empty object schema. - let input_schema = match &tool.parameters { - Some(p) => { - let mut map = serde_json::Map::new(); - map.insert("type".to_string(), Value::String(p._type.clone())); - map.insert( - "properties".to_string(), - Value::Object(p.properties.clone().into_iter().collect()), - ); - map.insert( - "required".to_string(), - Value::Array( - p.required - .iter() - .map(|s| Value::String(s.clone())) - .collect(), - ), - ); - Value::Object(map) - } - None => serde_json::json!({ "type": "object", "properties": {} }), - }; - + // Anthropic accepts standard JSON Schema unchanged. ClaudeTool { name: tool.name.clone(), description: tool.description.clone(), - input_schema, + input_schema: tool.parameters_json_schema(), } } } diff --git a/src/gemini/adapter.rs b/src/providers/gemini/adapter.rs similarity index 87% rename from src/gemini/adapter.rs rename to src/providers/gemini/adapter.rs index fe3cef8..8e557e0 100644 --- a/src/gemini/adapter.rs +++ b/src/providers/gemini/adapter.rs @@ -1,12 +1,18 @@ use crate::{ - client::{Completion, FunctionCall, MessageType, ModelRequest, Role, StreamEvent, Usage}, - gemini::types::{ - Content, FunctionCallPart, FunctionResponsePart, GeminiRequest, GeminiResponse, GeminiTool, - GeminiTools, GenerationConfig, Part, SystemInstructionContent, ThinkingConfig, - }, - provider::{BoxError, ProviderAdapter}, + error::LlmError, + provider::ProviderAdapter, + request::ModelRequest, + types::{Completion, FunctionCall, MessageType, Role, StreamEvent, Usage}, }; +use super::types::{ + Content, FunctionCallPart, FunctionResponsePart, GeminiRequest, GeminiResponse, GeminiTool, + GeminiTools, GenerationConfig, Part, SystemInstructionContent, ThinkingConfig, +}; + +/// Pure conversions between the common request/response types and the +/// Gemini `generateContent` wire format. No I/O. Shared by the direct API +/// and Vertex AI clients (they differ only in transport). #[derive(Debug, Default, Clone, Copy)] pub struct GeminiAdapter; @@ -78,12 +84,12 @@ impl ProviderAdapter for GeminiAdapter { } } - fn parse_completion(&self, body: &[u8]) -> Result { + fn parse_completion(&self, body: &[u8]) -> Result { let response_body: GeminiResponse = serde_json::from_slice(body)?; let content = response_body .get_text() - .ok_or_else(|| -> BoxError { "Missing completion from response".into() })?; + .ok_or_else(|| -> LlmError { "Missing completion from response".into() })?; Ok(Completion { completion: content, diff --git a/src/gemini/gcloud_helpers.rs b/src/providers/gemini/gcloud_helpers.rs similarity index 100% rename from src/gemini/gcloud_helpers.rs rename to src/providers/gemini/gcloud_helpers.rs diff --git a/src/gemini/mod.rs b/src/providers/gemini/mod.rs similarity index 92% rename from src/gemini/mod.rs rename to src/providers/gemini/mod.rs index 00b7820..97f7719 100644 --- a/src/gemini/mod.rs +++ b/src/providers/gemini/mod.rs @@ -7,7 +7,8 @@ mod tests; use async_trait::async_trait; -use crate::provider::{Action, BoxError, LlmClient, Transport}; +use crate::error::LlmError; +use crate::provider::{Action, LlmClient, Transport}; use gcloud_helpers::get_access_token; pub use adapter::GeminiAdapter; @@ -76,7 +77,7 @@ impl Transport for GeminiApiTransport { model: &str, action: Action, body: serde_json::Value, - ) -> Result { + ) -> Result { let url = format!( "https://generativelanguage.googleapis.com/v1beta/models/{}:{}", model, @@ -105,14 +106,14 @@ impl Transport for VertexTransport { model: &str, action: Action, body: serde_json::Value, - ) -> Result { + ) -> Result { let url = format!( "https://aiplatform.googleapis.com/v1/projects/{}/locations/global/publishers/google/models/{}:{}", self.project_name, model, method(action) ); - let access_token = get_access_token().await?; + let access_token = get_access_token().await.map_err(LlmError::Auth)?; Ok(self .client .post(url) diff --git a/src/gemini/tests.rs b/src/providers/gemini/tests.rs similarity index 97% rename from src/gemini/tests.rs rename to src/providers/gemini/tests.rs index 336ed0e..cb5367b 100644 --- a/src/gemini/tests.rs +++ b/src/providers/gemini/tests.rs @@ -5,8 +5,9 @@ use schemars::JsonSchema; use serde::{Deserialize, Serialize}; use crate::{ - client::{Message, Model, Settings, StreamEvent, Tool, Usage}, - gemini::{ + request::Model, + types::{Message, Settings, StreamEvent, Tool, Usage}, + providers::gemini::{ GeminiApiModel, GeminiVertexModel, adapter::GeminiAdapter, types::{GeminiModel, GeminiTool}, @@ -353,7 +354,7 @@ fn response_deserializes_when_content_has_no_parts() { // Some Gemini 3.x responses (e.g. thinking-only turns, MAX_TOKENS, safety // stops) return a candidate whose `content` has no `parts` field at all. // We must not fail to decode in that case. - use crate::gemini::types::GeminiResponse; + use crate::providers::gemini::types::GeminiResponse; let raw = r#"{ "candidates": [ @@ -383,7 +384,7 @@ fn response_with_partial_usage_metadata_reports_missing_counts_as_none() { // `promptTokenCount` in `usageMetadata`. Make sure the getters return // `None` for the missing counts (callers default them to 0) and the // response still decodes. - use crate::gemini::types::GeminiResponse; + use crate::providers::gemini::types::GeminiResponse; let raw = r#"{ "candidates": [ @@ -414,7 +415,7 @@ fn response_with_partial_usage_metadata_reports_missing_counts_as_none() { #[test] fn response_deserializes_when_candidate_has_no_content() { - use crate::gemini::types::GeminiResponse; + use crate::providers::gemini::types::GeminiResponse; let raw = r#"{ "candidates": [ @@ -432,12 +433,12 @@ fn make_direct_dummy(model: GeminiModel) -> GeminiApiModel { GeminiApiModel::new("dummy", model) } -fn build_body(request: crate::client::ModelRequest) -> crate::gemini::types::GeminiRequest { +fn build_body(request: crate::request::ModelRequest) -> super::types::GeminiRequest { GeminiAdapter.build_body(&request, "gemini-test", false) } -fn request_with_thinking(thinking_budget: Option) -> crate::client::ModelRequest { - crate::client::ModelRequest { +fn request_with_thinking(thinking_budget: Option) -> crate::request::ModelRequest { + crate::request::ModelRequest { system: None, messages: Some(vec![Message::user("hi".to_string())]), settings: Some(Settings { @@ -470,7 +471,7 @@ fn thinking_config_omitted_when_budget_is_none() { #[test] fn thinking_config_omitted_when_settings_is_none() { - let req = crate::client::ModelRequest { + let req = crate::request::ModelRequest { system: None, messages: Some(vec![Message::user("hi".to_string())]), settings: None, diff --git a/src/gemini/types.rs b/src/providers/gemini/types.rs similarity index 99% rename from src/gemini/types.rs rename to src/providers/gemini/types.rs index f2cdcf4..62869ca 100644 --- a/src/gemini/types.rs +++ b/src/providers/gemini/types.rs @@ -1,6 +1,6 @@ use std::collections::HashMap; -use crate::client::{Role, Tool}; +use crate::types::{Role, Tool}; use serde::{Deserialize, Serialize}; use serde_json::Value; diff --git a/src/providers/mod.rs b/src/providers/mod.rs new file mode 100644 index 0000000..440cf79 --- /dev/null +++ b/src/providers/mod.rs @@ -0,0 +1,3 @@ +pub mod claude; +pub mod gemini; +pub mod openai; diff --git a/src/openai/adapter.rs b/src/providers/openai/adapter.rs similarity index 94% rename from src/openai/adapter.rs rename to src/providers/openai/adapter.rs index aab8446..d93fb31 100644 --- a/src/openai/adapter.rs +++ b/src/providers/openai/adapter.rs @@ -1,12 +1,15 @@ use std::collections::HashMap; use crate::{ - client::{Completion, FunctionCall, MessageType, ModelRequest, StreamEvent, Usage}, - openai::types::{ - OpenAiInputItem, OpenAiRequest, OpenAiResponse, OpenAiTool, ResponsesStreamEvent, - synth_call_id, - }, - provider::{BoxError, ProviderAdapter}, + error::LlmError, + provider::ProviderAdapter, + request::ModelRequest, + types::{Completion, FunctionCall, MessageType, StreamEvent, Usage}, +}; + +use super::types::{ + OpenAiInputItem, OpenAiRequest, OpenAiResponse, OpenAiTool, ResponsesStreamEvent, + synth_call_id, }; #[derive(Debug, Default, Clone, Copy)] @@ -29,7 +32,7 @@ impl ProviderAdapter for OpenAiAdapter { for m in request.messages.clone().unwrap_or_default().iter() { match &m.message_type { MessageType::Text => match m.role { - Some(crate::client::Role::Model) => { + Some(crate::types::Role::Model) => { input.push(OpenAiInputItem::Message { role: "assistant".to_string(), content: m.content.clone(), @@ -79,7 +82,7 @@ impl ProviderAdapter for OpenAiAdapter { } } - fn parse_completion(&self, body: &[u8]) -> Result { + fn parse_completion(&self, body: &[u8]) -> Result { let body: OpenAiResponse = serde_json::from_slice(body)?; let text = body.get_text(); diff --git a/src/openai/mod.rs b/src/providers/openai/mod.rs similarity index 90% rename from src/openai/mod.rs rename to src/providers/openai/mod.rs index 8063ad0..adb143b 100644 --- a/src/openai/mod.rs +++ b/src/providers/openai/mod.rs @@ -6,7 +6,8 @@ mod tests; use async_trait::async_trait; -use crate::provider::{Action, BoxError, LlmClient, Transport}; +use crate::error::LlmError; +use crate::provider::{Action, LlmClient, Transport}; pub use adapter::OpenAiAdapter; pub use types::OpenAiModel; @@ -45,7 +46,7 @@ impl Transport for OpenAiTransport { _model: &str, _action: Action, body: serde_json::Value, - ) -> Result { + ) -> Result { Ok(self .client .post("https://api.openai.com/v1/responses") diff --git a/src/openai/tests.rs b/src/providers/openai/tests.rs similarity index 99% rename from src/openai/tests.rs rename to src/providers/openai/tests.rs index 7bc971c..047e628 100644 --- a/src/openai/tests.rs +++ b/src/providers/openai/tests.rs @@ -5,8 +5,9 @@ use schemars::JsonSchema; use serde::{Deserialize, Serialize}; use crate::{ - client::{Message, Model, Settings, StreamEvent, Tool, Usage}, - openai::{ + request::Model, + types::{Message, Settings, StreamEvent, Tool, Usage}, + providers::openai::{ OpenAiApiModel, types::{OpenAiModel, OpenAiTool}, }, diff --git a/src/openai/types.rs b/src/providers/openai/types.rs similarity index 85% rename from src/openai/types.rs rename to src/providers/openai/types.rs index f9277a7..d5e4a1a 100644 --- a/src/openai/types.rs +++ b/src/providers/openai/types.rs @@ -1,6 +1,6 @@ use std::collections::HashMap; -use crate::client::Tool; +use crate::types::Tool; use serde::{Deserialize, Serialize}; use serde_json::Value; @@ -66,33 +66,12 @@ pub struct OpenAiTool { impl OpenAiTool { pub fn from_tool(tool: &Tool) -> OpenAiTool { - let parameters = match &tool.parameters { - Some(p) => { - let mut map = serde_json::Map::new(); - map.insert("type".to_string(), Value::String(p._type.clone())); - map.insert( - "properties".to_string(), - Value::Object(p.properties.clone().into_iter().collect()), - ); - map.insert( - "required".to_string(), - Value::Array( - p.required - .iter() - .map(|s| Value::String(s.clone())) - .collect(), - ), - ); - Value::Object(map) - } - None => serde_json::json!({ "type": "object", "properties": {} }), - }; - + // OpenAI accepts standard JSON Schema unchanged. OpenAiTool { kind: "function", name: tool.name.clone(), description: tool.description.clone(), - parameters, + parameters: tool.parameters_json_schema(), strict: false, } } diff --git a/src/request.rs b/src/request.rs new file mode 100644 index 0000000..8788521 --- /dev/null +++ b/src/request.rs @@ -0,0 +1,97 @@ +//! The [`Model`] trait and the fluent [`ModelRequestBuilder`]. + +use async_trait::async_trait; + +use crate::error::LlmError; +use crate::types::{Completion, Message, Settings, StreamResult, Tool}; + +#[async_trait] +pub trait Model: Send + Sync { + async fn completion(&self, request: ModelRequest) -> Result; + + async fn stream_completion(&self, request: ModelRequest) -> Result; + + fn new_request(&self) -> ModelRequestBuilder<'_> + where + Self: Sized, + { + ModelRequestBuilder::new(self as &dyn Model) + } + + fn model_name(&self) -> String; +} + +pub struct ModelRequest { + pub system: Option, + pub messages: Option>, + pub settings: Option, + pub tools: Option>, +} + +#[derive(Clone)] +pub struct ModelRequestBuilder<'a> { + pub model: &'a dyn Model, + pub system: Option, + pub messages: Option>, + pub settings: Option, + pub tools: Option>, +} + +impl<'a> ModelRequestBuilder<'a> { + pub fn new(model: &'a dyn Model) -> Self { + ModelRequestBuilder { + model, + system: None, + messages: None, + settings: None, + tools: None, + } + } + + pub fn with_system(&mut self, system: String) -> &mut Self { + self.system = Some(system); + self + } + + pub fn with_message(&mut self, message: Message) -> &mut Self { + self.messages.get_or_insert_with(Vec::new).push(message); + self + } + + pub fn with_messages(&mut self, messages: Vec) -> &mut Self { + self.messages.get_or_insert_with(Vec::new).extend(messages); + self + } + + pub fn with_settings(&mut self, settings: Settings) -> &mut Self { + self.settings = Some(settings); + self + } + + pub fn with_tool(&mut self, tool: Tool) -> &mut Self { + self.tools.get_or_insert_with(Vec::new).push(tool); + self + } + + pub fn with_tools(&mut self, tools: Vec) -> &mut Self { + self.tools.get_or_insert_with(Vec::new).extend(tools); + self + } + + pub async fn completion(&self) -> Result { + self.model.completion(self.to_model_request()).await + } + + pub async fn stream(&self) -> Result { + self.model.stream_completion(self.to_model_request()).await + } + + pub fn to_model_request(&self) -> ModelRequest { + ModelRequest { + system: self.system.clone(), + messages: self.messages.clone(), + settings: self.settings.clone(), + tools: self.tools.clone(), + } + } +} diff --git a/src/client/tests.rs b/src/tests.rs similarity index 98% rename from src/client/tests.rs rename to src/tests.rs index 8885ed1..a408092 100644 --- a/src/client/tests.rs +++ b/src/tests.rs @@ -1,4 +1,6 @@ use super::*; +use serde_json::Value; +use std::collections::HashMap; use async_trait::async_trait; use schemars::JsonSchema; @@ -9,7 +11,7 @@ impl Model for MockModel { async fn completion( &self, _request: ModelRequest, - ) -> Result> { + ) -> Result { Ok(Completion { completion: "test".to_string(), usage: Usage { @@ -24,7 +26,7 @@ impl Model for MockModel { async fn stream_completion( &self, _request: ModelRequest, - ) -> Result> { + ) -> Result { use futures::stream; Ok(Box::pin(stream::iter(vec![ StreamEvent::Delta("test".to_string()), diff --git a/src/client/mod.rs b/src/types.rs similarity index 52% rename from src/client/mod.rs rename to src/types.rs index ddabbc8..897781f 100644 --- a/src/client/mod.rs +++ b/src/types.rs @@ -1,13 +1,12 @@ -use schemars::{JsonSchema, schema_for}; -use serde_json::{self, Value}; -use std::{collections::HashMap, error::Error, pin::Pin}; +//! Provider-agnostic data types: messages, tools, settings, usage and +//! completion/stream results. + +use std::{collections::HashMap, pin::Pin}; -use async_trait::async_trait; use futures::Stream; +use schemars::{JsonSchema, schema_for}; use serde::{Deserialize, Serialize}; - -#[cfg(test)] -mod tests; +use serde_json::{self, Value}; #[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] pub struct FunctionCall { @@ -69,7 +68,7 @@ pub struct Message { impl Message { pub fn user(content: String) -> Message { Message { - content: content, + content, role: Some(Role::User), message_type: MessageType::Text, } @@ -77,7 +76,7 @@ impl Message { pub fn model(content: String) -> Message { Message { - content: content, + content, role: Some(Role::Model), message_type: MessageType::Text, } @@ -107,28 +106,6 @@ impl Message { } } -#[async_trait] -pub trait Model: Send + Sync { - async fn completion( - &self, - request: ModelRequest, - ) -> Result>; - - async fn stream_completion( - &self, - request: ModelRequest, - ) -> Result>; - - fn new_request(&self) -> ModelRequestBuilder<'_> - where - Self: Sized, - { - ModelRequestBuilder::new(self as &dyn Model) - } - - fn model_name(&self) -> String; -} - #[derive(Clone)] pub struct Settings { pub max_tokens: Option, @@ -141,18 +118,31 @@ pub struct Settings { pub struct ToolParameters { #[serde(rename = "type")] pub _type: String, - #[serde(default = "default_properties")] + #[serde(default)] pub properties: HashMap, // TODO Eventually improve the typing here - #[serde(default = "default_required")] + #[serde(default)] pub required: Vec, } -fn default_properties() -> HashMap { - HashMap::new() -} +impl ToolParameters { + /// Standard JSON Schema object form: + /// `{ "type": ..., "properties": ..., "required": ... }`. + /// + /// Used verbatim by providers that accept plain JSON Schema (Anthropic, + /// OpenAI). Gemini applies its own conversion on top (uppercase type + /// names, `nullable`). + pub fn to_json_schema(&self) -> Value { + serde_json::json!({ + "type": self._type, + "properties": self.properties, + "required": self.required, + }) + } -fn default_required() -> Vec { - vec![] + /// Schema used when a tool declares no parameters: an empty object. + pub fn empty_json_schema() -> Value { + serde_json::json!({ "type": "object", "properties": {} }) + } } #[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] @@ -181,95 +171,13 @@ impl Tool { parameters: Some(parameters), }) } -} - -#[derive(Clone)] -pub struct ModelRequestBuilder<'a> { - pub model: &'a dyn Model, - pub system: Option, - pub messages: Option>, - pub settings: Option, - pub tools: Option>, -} - -pub struct ModelRequest { - pub system: Option, - pub messages: Option>, - pub settings: Option, - pub tools: Option>, -} - -impl<'a> ModelRequestBuilder<'a> { - pub fn new(model: &'a dyn Model) -> Self { - ModelRequestBuilder { - model, - system: None, - messages: None, - settings: None, - tools: None, - } - } - pub fn with_system(&mut self, system: String) -> &mut Self { - self.system = Some(system); - return self; - } - - pub fn with_message(&mut self, message: Message) -> &mut Self { - match &mut self.messages { - None => self.messages = Some(vec![message]), - Some(ms) => ms.push(message), - } - return self; - } - - pub fn with_messages(&mut self, messages: Vec) -> &mut Self { - match &mut self.messages { - None => self.messages = Some(messages), - Some(ms) => ms.extend(messages), - } - return self; - } - - pub fn with_settings(&mut self, settings: Settings) -> &mut Self { - self.settings = Some(settings); - return self; - } - - pub fn with_tool(&mut self, tool: Tool) -> &mut Self { - match self.tools { - None => self.tools = Some(vec![tool]), - Some(_) => { - self.tools.get_or_insert_with(Vec::new).push(tool); - } - } - return self; - } - - pub fn with_tools(&mut self, tools: Vec) -> &mut Self { - match self.tools { - None => self.tools = Some(tools), - Some(_) => { - self.tools.get_or_insert_with(Vec::new).extend(tools); - } - } - return self; - } - - pub async fn completion(&self) -> Result> { - self.model.completion(self.to_model_request()).await - } - - pub async fn stream(&self) -> Result> { - self.model.stream_completion(self.to_model_request()).await - } - - pub fn to_model_request(&self) -> ModelRequest { - ModelRequest { - system: self.system.clone(), - messages: self.messages.clone(), - settings: self.settings.clone(), - tools: self.tools.clone(), - } + /// The tool's parameters as standard JSON Schema, or an empty object + /// schema when the tool takes no parameters. + pub fn parameters_json_schema(&self) -> Value { + self.parameters + .as_ref() + .map(ToolParameters::to_json_schema) + .unwrap_or_else(ToolParameters::empty_json_schema) } }