mirror of
https://github.com/outbackdingo/optimclaw.git
synced 2026-09-03 01:59:23 +00:00
Simplify workspace to path-based storage, remove legacy code
- Consolidate all migrations into V1__initial.sql - Replace DocType enum with flexible path-based file storage - Add list_workspace_files SQL function for directory listing - Update memory tools for path-based API (memory_read, memory_write, memory_search, memory_list) - Remove unused OpenAI/Anthropic providers (NEAR AI only) - Simplify config to remove multi-provider support - Update CLAUDE.md documentation Co-Authored-By: Claude Opus 4.5 <[email protected]>
This commit is contained in:
co-authored by
Claude Opus 4.5
parent
f29892b3fb
commit
3718cfa767
@@ -1,348 +0,0 @@
|
||||
//! Anthropic LLM provider implementation.
|
||||
|
||||
use async_trait::async_trait;
|
||||
use reqwest::Client;
|
||||
use rust_decimal::Decimal;
|
||||
use rust_decimal_macros::dec;
|
||||
use secrecy::ExposeSecret;
|
||||
use serde::{Deserialize, Serialize};
|
||||
|
||||
use crate::config::AnthropicConfig;
|
||||
use crate::error::LlmError;
|
||||
use crate::llm::provider::{
|
||||
ChatMessage, CompletionRequest, CompletionResponse, FinishReason, LlmProvider, Role, ToolCall,
|
||||
ToolCompletionRequest, ToolCompletionResponse,
|
||||
};
|
||||
|
||||
/// Anthropic API provider.
|
||||
pub struct AnthropicProvider {
|
||||
client: Client,
|
||||
config: AnthropicConfig,
|
||||
base_url: String,
|
||||
}
|
||||
|
||||
impl AnthropicProvider {
|
||||
/// Create a new Anthropic provider.
|
||||
pub fn new(config: AnthropicConfig) -> Self {
|
||||
let base_url = config
|
||||
.base_url
|
||||
.clone()
|
||||
.unwrap_or_else(|| "https://api.anthropic.com/v1".to_string());
|
||||
|
||||
Self {
|
||||
client: Client::new(),
|
||||
config,
|
||||
base_url,
|
||||
}
|
||||
}
|
||||
|
||||
fn build_messages(&self, messages: &[ChatMessage]) -> (Option<String>, Vec<AnthropicMessage>) {
|
||||
let mut system_message = None;
|
||||
let mut anthropic_messages = Vec::new();
|
||||
|
||||
for msg in messages {
|
||||
match msg.role {
|
||||
Role::System => {
|
||||
// Anthropic uses a separate system parameter
|
||||
system_message = Some(msg.content.clone());
|
||||
}
|
||||
Role::User => {
|
||||
anthropic_messages.push(AnthropicMessage {
|
||||
role: "user".to_string(),
|
||||
content: AnthropicContent::Text(msg.content.clone()),
|
||||
});
|
||||
}
|
||||
Role::Assistant => {
|
||||
anthropic_messages.push(AnthropicMessage {
|
||||
role: "assistant".to_string(),
|
||||
content: AnthropicContent::Text(msg.content.clone()),
|
||||
});
|
||||
}
|
||||
Role::Tool => {
|
||||
// Tool results in Anthropic format
|
||||
anthropic_messages.push(AnthropicMessage {
|
||||
role: "user".to_string(),
|
||||
content: AnthropicContent::ToolResult {
|
||||
tool_use_id: msg.tool_call_id.clone().unwrap_or_default(),
|
||||
content: msg.content.clone(),
|
||||
},
|
||||
});
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
(system_message, anthropic_messages)
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Serialize)]
|
||||
struct AnthropicRequest {
|
||||
model: String,
|
||||
messages: Vec<AnthropicMessage>,
|
||||
max_tokens: u32,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
system: Option<String>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
temperature: Option<f32>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
tools: Option<Vec<AnthropicTool>>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
tool_choice: Option<AnthropicToolChoice>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Serialize)]
|
||||
struct AnthropicMessage {
|
||||
role: String,
|
||||
content: AnthropicContent,
|
||||
}
|
||||
|
||||
#[derive(Debug, Serialize)]
|
||||
#[serde(untagged)]
|
||||
enum AnthropicContent {
|
||||
Text(String),
|
||||
#[serde(rename_all = "snake_case")]
|
||||
ToolResult {
|
||||
#[serde(rename = "type")]
|
||||
tool_use_id: String,
|
||||
content: String,
|
||||
},
|
||||
Blocks(Vec<AnthropicContentBlock>),
|
||||
}
|
||||
|
||||
#[derive(Debug, Serialize, Deserialize)]
|
||||
#[serde(tag = "type")]
|
||||
enum AnthropicContentBlock {
|
||||
#[serde(rename = "text")]
|
||||
Text { text: String },
|
||||
#[serde(rename = "tool_use")]
|
||||
ToolUse {
|
||||
id: String,
|
||||
name: String,
|
||||
input: serde_json::Value,
|
||||
},
|
||||
#[serde(rename = "tool_result")]
|
||||
ToolResult {
|
||||
tool_use_id: String,
|
||||
content: String,
|
||||
},
|
||||
}
|
||||
|
||||
#[derive(Debug, Serialize)]
|
||||
struct AnthropicTool {
|
||||
name: String,
|
||||
description: String,
|
||||
input_schema: serde_json::Value,
|
||||
}
|
||||
|
||||
#[derive(Debug, Serialize)]
|
||||
struct AnthropicToolChoice {
|
||||
#[serde(rename = "type")]
|
||||
choice_type: String,
|
||||
}
|
||||
|
||||
#[derive(Debug, Deserialize)]
|
||||
struct AnthropicResponse {
|
||||
content: Vec<AnthropicContentBlock>,
|
||||
stop_reason: Option<String>,
|
||||
usage: AnthropicUsage,
|
||||
}
|
||||
|
||||
#[derive(Debug, Deserialize)]
|
||||
struct AnthropicUsage {
|
||||
input_tokens: u32,
|
||||
output_tokens: u32,
|
||||
}
|
||||
|
||||
#[derive(Debug, Deserialize)]
|
||||
struct AnthropicError {
|
||||
error: AnthropicErrorDetail,
|
||||
}
|
||||
|
||||
#[derive(Debug, Deserialize)]
|
||||
struct AnthropicErrorDetail {
|
||||
message: String,
|
||||
#[serde(rename = "type")]
|
||||
error_type: String,
|
||||
}
|
||||
|
||||
fn parse_finish_reason(reason: Option<&str>) -> FinishReason {
|
||||
match reason {
|
||||
Some("end_turn") | Some("stop_sequence") => FinishReason::Stop,
|
||||
Some("max_tokens") => FinishReason::Length,
|
||||
Some("tool_use") => FinishReason::ToolUse,
|
||||
_ => FinishReason::Unknown,
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl LlmProvider for AnthropicProvider {
|
||||
fn model_name(&self) -> &str {
|
||||
&self.config.model
|
||||
}
|
||||
|
||||
fn cost_per_token(&self) -> (Decimal, Decimal) {
|
||||
// Pricing for Claude models (per 1M tokens, converted to per token)
|
||||
match self.config.model.as_str() {
|
||||
m if m.contains("opus") => {
|
||||
(dec!(0.000015), dec!(0.000075)) // $15/$75 per 1M
|
||||
}
|
||||
m if m.contains("sonnet") => {
|
||||
(dec!(0.000003), dec!(0.000015)) // $3/$15 per 1M
|
||||
}
|
||||
m if m.contains("haiku") => {
|
||||
(dec!(0.00000025), dec!(0.00000125)) // $0.25/$1.25 per 1M
|
||||
}
|
||||
_ => (dec!(0.000003), dec!(0.000015)), // Default to Sonnet pricing
|
||||
}
|
||||
}
|
||||
|
||||
async fn complete(&self, request: CompletionRequest) -> Result<CompletionResponse, LlmError> {
|
||||
let (system, messages) = self.build_messages(&request.messages);
|
||||
|
||||
let anthropic_request = AnthropicRequest {
|
||||
model: self.config.model.clone(),
|
||||
messages,
|
||||
max_tokens: request.max_tokens.unwrap_or(4096),
|
||||
system,
|
||||
temperature: request.temperature,
|
||||
tools: None,
|
||||
tool_choice: None,
|
||||
};
|
||||
|
||||
let response = self
|
||||
.client
|
||||
.post(format!("{}/messages", self.base_url))
|
||||
.header("x-api-key", self.config.api_key.expose_secret())
|
||||
.header("anthropic-version", "2023-06-01")
|
||||
.header("Content-Type", "application/json")
|
||||
.json(&anthropic_request)
|
||||
.send()
|
||||
.await?;
|
||||
|
||||
if !response.status().is_success() {
|
||||
let error: AnthropicError =
|
||||
response
|
||||
.json()
|
||||
.await
|
||||
.map_err(|e| LlmError::InvalidResponse {
|
||||
provider: "anthropic".to_string(),
|
||||
reason: format!("Failed to parse error response: {}", e),
|
||||
})?;
|
||||
return Err(LlmError::RequestFailed {
|
||||
provider: "anthropic".to_string(),
|
||||
reason: error.error.message,
|
||||
});
|
||||
}
|
||||
|
||||
let anthropic_response: AnthropicResponse = response.json().await?;
|
||||
|
||||
// Extract text content
|
||||
let content = anthropic_response
|
||||
.content
|
||||
.iter()
|
||||
.filter_map(|block| match block {
|
||||
AnthropicContentBlock::Text { text } => Some(text.clone()),
|
||||
_ => None,
|
||||
})
|
||||
.collect::<Vec<_>>()
|
||||
.join("\n");
|
||||
|
||||
Ok(CompletionResponse {
|
||||
content,
|
||||
input_tokens: anthropic_response.usage.input_tokens,
|
||||
output_tokens: anthropic_response.usage.output_tokens,
|
||||
finish_reason: parse_finish_reason(anthropic_response.stop_reason.as_deref()),
|
||||
})
|
||||
}
|
||||
|
||||
async fn complete_with_tools(
|
||||
&self,
|
||||
request: ToolCompletionRequest,
|
||||
) -> Result<ToolCompletionResponse, LlmError> {
|
||||
let (system, messages) = self.build_messages(&request.messages);
|
||||
|
||||
let tools: Vec<AnthropicTool> = request
|
||||
.tools
|
||||
.iter()
|
||||
.map(|t| AnthropicTool {
|
||||
name: t.name.clone(),
|
||||
description: t.description.clone(),
|
||||
input_schema: t.parameters.clone(),
|
||||
})
|
||||
.collect();
|
||||
|
||||
let tool_choice = request.tool_choice.as_ref().map(|c| AnthropicToolChoice {
|
||||
choice_type: match c.as_str() {
|
||||
"auto" => "auto".to_string(),
|
||||
"required" => "any".to_string(),
|
||||
"none" => "none".to_string(),
|
||||
_ => "auto".to_string(),
|
||||
},
|
||||
});
|
||||
|
||||
let anthropic_request = AnthropicRequest {
|
||||
model: self.config.model.clone(),
|
||||
messages,
|
||||
max_tokens: request.max_tokens.unwrap_or(4096),
|
||||
system,
|
||||
temperature: None,
|
||||
tools: Some(tools),
|
||||
tool_choice,
|
||||
};
|
||||
|
||||
let response = self
|
||||
.client
|
||||
.post(format!("{}/messages", self.base_url))
|
||||
.header("x-api-key", self.config.api_key.expose_secret())
|
||||
.header("anthropic-version", "2023-06-01")
|
||||
.header("Content-Type", "application/json")
|
||||
.json(&anthropic_request)
|
||||
.send()
|
||||
.await?;
|
||||
|
||||
if !response.status().is_success() {
|
||||
let error: AnthropicError =
|
||||
response
|
||||
.json()
|
||||
.await
|
||||
.map_err(|e| LlmError::InvalidResponse {
|
||||
provider: "anthropic".to_string(),
|
||||
reason: format!("Failed to parse error response: {}", e),
|
||||
})?;
|
||||
return Err(LlmError::RequestFailed {
|
||||
provider: "anthropic".to_string(),
|
||||
reason: error.error.message,
|
||||
});
|
||||
}
|
||||
|
||||
let anthropic_response: AnthropicResponse = response.json().await?;
|
||||
|
||||
// Extract text and tool calls
|
||||
let mut content = None;
|
||||
let mut tool_calls = Vec::new();
|
||||
|
||||
for block in anthropic_response.content {
|
||||
match block {
|
||||
AnthropicContentBlock::Text { text } => {
|
||||
content = Some(text);
|
||||
}
|
||||
AnthropicContentBlock::ToolUse { id, name, input } => {
|
||||
tool_calls.push(ToolCall {
|
||||
id,
|
||||
name,
|
||||
arguments: input,
|
||||
});
|
||||
}
|
||||
_ => {}
|
||||
}
|
||||
}
|
||||
|
||||
Ok(ToolCompletionResponse {
|
||||
content,
|
||||
tool_calls,
|
||||
input_tokens: anthropic_response.usage.input_tokens,
|
||||
output_tokens: anthropic_response.usage.output_tokens,
|
||||
finish_reason: parse_finish_reason(anthropic_response.stop_reason.as_deref()),
|
||||
})
|
||||
}
|
||||
}
|
||||
+3
-31
@@ -1,17 +1,12 @@
|
||||
//! LLM integration for the agent.
|
||||
//!
|
||||
//! Provides a unified interface to different LLM providers (OpenAI, Anthropic, NEAR AI)
|
||||
//! and implements reasoning capabilities for planning, tool selection, and evaluation.
|
||||
//! Uses the NEAR AI chat-api as the unified LLM provider.
|
||||
|
||||
mod anthropic;
|
||||
mod nearai;
|
||||
mod openai;
|
||||
mod provider;
|
||||
mod reasoning;
|
||||
|
||||
pub use anthropic::AnthropicProvider;
|
||||
pub use nearai::NearAiProvider;
|
||||
pub use openai::OpenAiProvider;
|
||||
pub use provider::{
|
||||
ChatMessage, CompletionRequest, CompletionResponse, LlmProvider, Role, ToolCall,
|
||||
ToolCompletionRequest, ToolCompletionResponse, ToolDefinition, ToolResult,
|
||||
@@ -20,33 +15,10 @@ pub use reasoning::{ActionPlan, Reasoning, ReasoningContext, ToolSelection};
|
||||
|
||||
use std::sync::Arc;
|
||||
|
||||
use crate::config::{LlmConfig, LlmProvider as LlmProviderType};
|
||||
use crate::config::LlmConfig;
|
||||
use crate::error::LlmError;
|
||||
|
||||
/// Create an LLM provider based on configuration.
|
||||
pub fn create_llm_provider(config: &LlmConfig) -> Result<Arc<dyn LlmProvider>, LlmError> {
|
||||
match config.provider {
|
||||
LlmProviderType::OpenAi => {
|
||||
let openai_config = config.openai.as_ref().ok_or_else(|| LlmError::AuthFailed {
|
||||
provider: "openai".to_string(),
|
||||
})?;
|
||||
Ok(Arc::new(OpenAiProvider::new(openai_config.clone())))
|
||||
}
|
||||
LlmProviderType::Anthropic => {
|
||||
let anthropic_config =
|
||||
config
|
||||
.anthropic
|
||||
.as_ref()
|
||||
.ok_or_else(|| LlmError::AuthFailed {
|
||||
provider: "anthropic".to_string(),
|
||||
})?;
|
||||
Ok(Arc::new(AnthropicProvider::new(anthropic_config.clone())))
|
||||
}
|
||||
LlmProviderType::NearAi => {
|
||||
let nearai_config = config.nearai.as_ref().ok_or_else(|| LlmError::AuthFailed {
|
||||
provider: "nearai".to_string(),
|
||||
})?;
|
||||
Ok(Arc::new(NearAiProvider::new(nearai_config.clone())))
|
||||
}
|
||||
}
|
||||
Ok(Arc::new(NearAiProvider::new(config.nearai.clone())))
|
||||
}
|
||||
|
||||
@@ -1,335 +0,0 @@
|
||||
//! OpenAI LLM provider implementation.
|
||||
|
||||
use async_trait::async_trait;
|
||||
use reqwest::Client;
|
||||
use rust_decimal::Decimal;
|
||||
use rust_decimal_macros::dec;
|
||||
use secrecy::ExposeSecret;
|
||||
use serde::{Deserialize, Serialize};
|
||||
|
||||
use crate::config::OpenAiConfig;
|
||||
use crate::error::LlmError;
|
||||
use crate::llm::provider::{
|
||||
ChatMessage, CompletionRequest, CompletionResponse, FinishReason, LlmProvider, Role, ToolCall,
|
||||
ToolCompletionRequest, ToolCompletionResponse,
|
||||
};
|
||||
|
||||
/// OpenAI API provider.
|
||||
pub struct OpenAiProvider {
|
||||
client: Client,
|
||||
config: OpenAiConfig,
|
||||
base_url: String,
|
||||
}
|
||||
|
||||
impl OpenAiProvider {
|
||||
/// Create a new OpenAI provider.
|
||||
pub fn new(config: OpenAiConfig) -> Self {
|
||||
let base_url = config
|
||||
.base_url
|
||||
.clone()
|
||||
.unwrap_or_else(|| "https://api.openai.com/v1".to_string());
|
||||
|
||||
Self {
|
||||
client: Client::new(),
|
||||
config,
|
||||
base_url,
|
||||
}
|
||||
}
|
||||
|
||||
fn build_messages(&self, messages: &[ChatMessage]) -> Vec<OpenAiMessage> {
|
||||
messages
|
||||
.iter()
|
||||
.map(|m| OpenAiMessage {
|
||||
role: match m.role {
|
||||
Role::System => "system".to_string(),
|
||||
Role::User => "user".to_string(),
|
||||
Role::Assistant => "assistant".to_string(),
|
||||
Role::Tool => "tool".to_string(),
|
||||
},
|
||||
content: Some(m.content.clone()),
|
||||
tool_call_id: m.tool_call_id.clone(),
|
||||
name: m.name.clone(),
|
||||
tool_calls: None,
|
||||
})
|
||||
.collect()
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Serialize)]
|
||||
struct OpenAiRequest {
|
||||
model: String,
|
||||
messages: Vec<OpenAiMessage>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
max_tokens: Option<u32>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
temperature: Option<f32>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
tools: Option<Vec<OpenAiTool>>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
tool_choice: Option<serde_json::Value>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Serialize, Deserialize)]
|
||||
struct OpenAiMessage {
|
||||
role: String,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
content: Option<String>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
tool_call_id: Option<String>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
name: Option<String>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
tool_calls: Option<Vec<OpenAiToolCall>>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Serialize)]
|
||||
struct OpenAiTool {
|
||||
#[serde(rename = "type")]
|
||||
tool_type: String,
|
||||
function: OpenAiFunction,
|
||||
}
|
||||
|
||||
#[derive(Debug, Serialize)]
|
||||
struct OpenAiFunction {
|
||||
name: String,
|
||||
description: String,
|
||||
parameters: serde_json::Value,
|
||||
}
|
||||
|
||||
#[derive(Debug, Deserialize)]
|
||||
struct OpenAiResponse {
|
||||
choices: Vec<OpenAiChoice>,
|
||||
usage: OpenAiUsage,
|
||||
}
|
||||
|
||||
#[derive(Debug, Deserialize)]
|
||||
struct OpenAiChoice {
|
||||
message: OpenAiResponseMessage,
|
||||
finish_reason: Option<String>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Deserialize)]
|
||||
struct OpenAiResponseMessage {
|
||||
content: Option<String>,
|
||||
tool_calls: Option<Vec<OpenAiToolCall>>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Serialize, Deserialize, Clone)]
|
||||
struct OpenAiToolCall {
|
||||
id: String,
|
||||
#[serde(rename = "type")]
|
||||
call_type: String,
|
||||
function: OpenAiFunctionCall,
|
||||
}
|
||||
|
||||
#[derive(Debug, Serialize, Deserialize, Clone)]
|
||||
struct OpenAiFunctionCall {
|
||||
name: String,
|
||||
arguments: String,
|
||||
}
|
||||
|
||||
#[derive(Debug, Deserialize)]
|
||||
struct OpenAiUsage {
|
||||
prompt_tokens: u32,
|
||||
completion_tokens: u32,
|
||||
}
|
||||
|
||||
#[derive(Debug, Deserialize)]
|
||||
struct OpenAiError {
|
||||
error: OpenAiErrorDetail,
|
||||
}
|
||||
|
||||
#[derive(Debug, Deserialize)]
|
||||
struct OpenAiErrorDetail {
|
||||
message: String,
|
||||
#[serde(rename = "type")]
|
||||
error_type: Option<String>,
|
||||
}
|
||||
|
||||
fn parse_finish_reason(reason: Option<&str>) -> FinishReason {
|
||||
match reason {
|
||||
Some("stop") => FinishReason::Stop,
|
||||
Some("length") => FinishReason::Length,
|
||||
Some("tool_calls") => FinishReason::ToolUse,
|
||||
Some("content_filter") => FinishReason::ContentFilter,
|
||||
_ => FinishReason::Unknown,
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl LlmProvider for OpenAiProvider {
|
||||
fn model_name(&self) -> &str {
|
||||
&self.config.model
|
||||
}
|
||||
|
||||
fn cost_per_token(&self) -> (Decimal, Decimal) {
|
||||
// Pricing for GPT-4 Turbo (per 1M tokens, converted to per token)
|
||||
// These are approximate and should be updated based on actual pricing
|
||||
match self.config.model.as_str() {
|
||||
m if m.contains("gpt-4-turbo") || m.contains("gpt-4o") => {
|
||||
(dec!(0.00001), dec!(0.00003)) // $10/$30 per 1M
|
||||
}
|
||||
m if m.contains("gpt-4") => {
|
||||
(dec!(0.00003), dec!(0.00006)) // $30/$60 per 1M
|
||||
}
|
||||
m if m.contains("gpt-3.5") => {
|
||||
(dec!(0.0000005), dec!(0.0000015)) // $0.50/$1.50 per 1M
|
||||
}
|
||||
_ => (dec!(0.00001), dec!(0.00003)), // Default to GPT-4 Turbo pricing
|
||||
}
|
||||
}
|
||||
|
||||
async fn complete(&self, request: CompletionRequest) -> Result<CompletionResponse, LlmError> {
|
||||
let openai_request = OpenAiRequest {
|
||||
model: self.config.model.clone(),
|
||||
messages: self.build_messages(&request.messages),
|
||||
max_tokens: request.max_tokens,
|
||||
temperature: request.temperature,
|
||||
tools: None,
|
||||
tool_choice: None,
|
||||
};
|
||||
|
||||
let response = self
|
||||
.client
|
||||
.post(format!("{}/chat/completions", self.base_url))
|
||||
.header(
|
||||
"Authorization",
|
||||
format!("Bearer {}", self.config.api_key.expose_secret()),
|
||||
)
|
||||
.header("Content-Type", "application/json")
|
||||
.json(&openai_request)
|
||||
.send()
|
||||
.await?;
|
||||
|
||||
if !response.status().is_success() {
|
||||
let error: OpenAiError =
|
||||
response
|
||||
.json()
|
||||
.await
|
||||
.map_err(|e| LlmError::InvalidResponse {
|
||||
provider: "openai".to_string(),
|
||||
reason: format!("Failed to parse error response: {}", e),
|
||||
})?;
|
||||
return Err(LlmError::RequestFailed {
|
||||
provider: "openai".to_string(),
|
||||
reason: error.error.message,
|
||||
});
|
||||
}
|
||||
|
||||
let openai_response: OpenAiResponse = response.json().await?;
|
||||
|
||||
let choice = openai_response
|
||||
.choices
|
||||
.first()
|
||||
.ok_or_else(|| LlmError::InvalidResponse {
|
||||
provider: "openai".to_string(),
|
||||
reason: "No choices in response".to_string(),
|
||||
})?;
|
||||
|
||||
Ok(CompletionResponse {
|
||||
content: choice.message.content.clone().unwrap_or_default(),
|
||||
input_tokens: openai_response.usage.prompt_tokens,
|
||||
output_tokens: openai_response.usage.completion_tokens,
|
||||
finish_reason: parse_finish_reason(choice.finish_reason.as_deref()),
|
||||
})
|
||||
}
|
||||
|
||||
async fn complete_with_tools(
|
||||
&self,
|
||||
request: ToolCompletionRequest,
|
||||
) -> Result<ToolCompletionResponse, LlmError> {
|
||||
let tools: Vec<OpenAiTool> = request
|
||||
.tools
|
||||
.iter()
|
||||
.map(|t| OpenAiTool {
|
||||
tool_type: "function".to_string(),
|
||||
function: OpenAiFunction {
|
||||
name: t.name.clone(),
|
||||
description: t.description.clone(),
|
||||
parameters: t.parameters.clone(),
|
||||
},
|
||||
})
|
||||
.collect();
|
||||
|
||||
let tool_choice = request.tool_choice.as_ref().map(|c| match c.as_str() {
|
||||
"auto" => serde_json::json!("auto"),
|
||||
"required" => serde_json::json!("required"),
|
||||
"none" => serde_json::json!("none"),
|
||||
_ => serde_json::json!("auto"),
|
||||
});
|
||||
|
||||
let openai_request = OpenAiRequest {
|
||||
model: self.config.model.clone(),
|
||||
messages: self.build_messages(&request.messages),
|
||||
max_tokens: request.max_tokens,
|
||||
temperature: None,
|
||||
tools: Some(tools),
|
||||
tool_choice,
|
||||
};
|
||||
|
||||
let response = self
|
||||
.client
|
||||
.post(format!("{}/chat/completions", self.base_url))
|
||||
.header(
|
||||
"Authorization",
|
||||
format!("Bearer {}", self.config.api_key.expose_secret()),
|
||||
)
|
||||
.header("Content-Type", "application/json")
|
||||
.json(&openai_request)
|
||||
.send()
|
||||
.await?;
|
||||
|
||||
if !response.status().is_success() {
|
||||
let error: OpenAiError =
|
||||
response
|
||||
.json()
|
||||
.await
|
||||
.map_err(|e| LlmError::InvalidResponse {
|
||||
provider: "openai".to_string(),
|
||||
reason: format!("Failed to parse error response: {}", e),
|
||||
})?;
|
||||
return Err(LlmError::RequestFailed {
|
||||
provider: "openai".to_string(),
|
||||
reason: error.error.message,
|
||||
});
|
||||
}
|
||||
|
||||
let openai_response: OpenAiResponse = response.json().await?;
|
||||
|
||||
let choice = openai_response
|
||||
.choices
|
||||
.first()
|
||||
.ok_or_else(|| LlmError::InvalidResponse {
|
||||
provider: "openai".to_string(),
|
||||
reason: "No choices in response".to_string(),
|
||||
})?;
|
||||
|
||||
let tool_calls: Vec<ToolCall> = choice
|
||||
.message
|
||||
.tool_calls
|
||||
.as_ref()
|
||||
.map(|calls| {
|
||||
calls
|
||||
.iter()
|
||||
.filter_map(|c| {
|
||||
let args: serde_json::Value =
|
||||
serde_json::from_str(&c.function.arguments).ok()?;
|
||||
Some(ToolCall {
|
||||
id: c.id.clone(),
|
||||
name: c.function.name.clone(),
|
||||
arguments: args,
|
||||
})
|
||||
})
|
||||
.collect()
|
||||
})
|
||||
.unwrap_or_default();
|
||||
|
||||
Ok(ToolCompletionResponse {
|
||||
content: choice.message.content.clone(),
|
||||
tool_calls,
|
||||
input_tokens: openai_response.usage.prompt_tokens,
|
||||
output_tokens: openai_response.usage.completion_tokens,
|
||||
finish_reason: parse_finish_reason(choice.finish_reason.as_deref()),
|
||||
})
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user