volition_core/providers/
openai.rs1use super::Provider;
3use crate::config::ModelConfig;
4use crate::models::chat::{ApiResponse, ChatMessage, Choice};
5use crate::models::tools::ToolDefinition;
6use anyhow::{Result, anyhow, Context};
7use async_trait::async_trait;
8use reqwest::Client;
9use serde_json::{json, Value};
10use tracing::{debug, warn};
11
12const DEFAULT_OPENAI_ENDPOINT: &str = "https://api.openai.com/v1/chat/completions";
13
14#[derive(Clone)]
15pub struct OpenAIProvider {
16 config: ModelConfig,
17 http_client: Client,
18 api_key: String,
19}
20
21impl OpenAIProvider {
22 pub fn new(config: ModelConfig, http_client: Client, api_key: String) -> Self {
23 Self {
24 config,
25 http_client,
26 api_key,
27 }
28 }
29
30 fn build_payload(
31 &self,
32 messages: Vec<ChatMessage>,
33 tools: Option<&[ToolDefinition]>,
34 ) -> Result<Value> {
35 debug!("Building OpenAI payload...");
36 debug!("Model name: {}", self.config.model_name);
37 debug!("Message count: {}", messages.len());
38
39 let mut payload = json!({
40 "model": self.config.model_name,
41 "messages": messages.iter().map(|msg| {
42 json!({
43 "role": msg.role,
44 "content": msg.content.as_deref().unwrap_or_default()
45 })
46 }).collect::<Vec<_>>()
47 });
48
49 if let Some(tools) = tools {
51 if !tools.is_empty() {
52 let functions: Vec<Value> = tools
53 .iter()
54 .map(|t| {
55 json!({
56 "name": t.name,
57 "description": t.description,
58 "parameters": t.parameters
59 })
60 })
61 .collect();
62 payload["functions"] = json!(functions);
63 payload["function_call"] = json!("auto");
64 }
65 }
66
67 if let Some(params) = &self.config.parameters {
69 if let Some(temperature) = params.get("temperature").and_then(|t| t.as_float()) {
70 payload["temperature"] = json!(temperature);
71 }
72 }
74
75 debug!("Final payload: {}", serde_json::to_string_pretty(&payload)?);
76 Ok(payload)
77 }
78
79 fn parse_response(&self, response_body: &str) -> Result<ApiResponse> {
80 debug!("Parsing OpenAI response...");
81 debug!("Response body: {}", response_body);
82
83 let raw_response: Value = serde_json::from_str(response_body)?;
84
85 let choice = &raw_response["choices"][0];
86 let message = &choice["message"];
87
88 let content = message["content"]
89 .as_str()
90 .ok_or_else(|| anyhow!("Missing content in OpenAI response"))?
91 .to_string();
92 debug!("Extracted content: {}", content);
93
94 let finish_reason = choice["finish_reason"]
95 .as_str()
96 .unwrap_or("stop")
97 .to_string();
98 debug!("Finish reason: {}", finish_reason);
99
100 let usage = &raw_response["usage"];
101 let prompt_tokens = usage["prompt_tokens"].as_u64().unwrap_or(0) as u32;
102 let completion_tokens = usage["completion_tokens"].as_u64().unwrap_or(0) as u32;
103 let total_tokens = usage["total_tokens"].as_u64().unwrap_or(0) as u32;
104 debug!("Token usage - prompt: {}, completion: {}, total: {}",
105 prompt_tokens, completion_tokens, total_tokens);
106
107 let mut tool_calls = None;
108 if let Some(function_call) = message.get("function_call") {
109 if let (Some(name), Some(arguments)) = (
110 function_call["name"].as_str(),
111 function_call["arguments"].as_str(),
112 ) {
113 tool_calls = Some(vec![crate::models::tools::ToolCall {
114 id: format!("call_{}", name),
115 call_type: "function".to_string(),
116 function: crate::models::tools::ToolFunction {
117 name: name.to_string(),
118 arguments: arguments.to_string(),
119 },
120 }]);
121 }
122 }
123
124 let result = ApiResponse {
125 id: raw_response["id"]
126 .as_str()
127 .map(|s| s.to_string())
128 .unwrap_or_default(),
129 content: content.clone(),
130 finish_reason: finish_reason.clone(),
131 prompt_tokens,
132 completion_tokens,
133 total_tokens,
134 choices: vec![Choice {
135 index: 0,
136 message: ChatMessage {
137 role: "assistant".to_string(),
138 content: Some(content),
139 tool_calls,
140 tool_call_id: None,
141 },
142 finish_reason,
143 }],
144 };
145
146 debug!("Parsed response: {:?}", result);
147 Ok(result)
148 }
149
150 async fn call_chat_completion_api(
151 &self,
152 messages: Vec<ChatMessage>,
153 tools: Option<&[ToolDefinition]>,
154 ) -> Result<ApiResponse> {
155 let endpoint = self.config.endpoint.as_deref().unwrap_or_else(|| {
156 warn!("No endpoint specified for OpenAI provider model {}, using default: {}", self.config.model_name, DEFAULT_OPENAI_ENDPOINT);
157 DEFAULT_OPENAI_ENDPOINT
158 });
159
160 if self.api_key.is_empty() {
161 warn!(
162 "API key is empty for OpenAI provider model {}. The API call will likely fail.",
163 self.config.model_name
164 );
165 }
166
167 let payload = self.build_payload(messages, tools)?;
168
169 let response = self
170 .http_client
171 .post(endpoint)
172 .header("Content-Type", "application/json")
173 .header("Authorization", format!("Bearer {}", self.api_key))
174 .json(&payload)
175 .send()
176 .await
177 .context("Failed to send request to OpenAI API")?;
178
179 let response_body = response
180 .text()
181 .await
182 .context("Failed to read response from OpenAI API")?;
183
184 self.parse_response(&response_body)
185 }
186}
187
188#[async_trait]
189impl Provider for OpenAIProvider {
190 fn name(&self) -> &str {
191 &self.config.model_name
192 }
193
194 async fn get_completion(
195 &self,
196 messages: Vec<ChatMessage>,
197 tools: Option<&[ToolDefinition]>,
198 ) -> Result<ApiResponse> {
199 self.call_chat_completion_api(messages, tools).await
200 }
201}