Skip to main content

volition_core/providers/
openai.rs

1// volition-agent-core/src/providers/openai.rs
2use 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        // Add tools if present
50        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        // Add parameters if present
68        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            // Add other OpenAI-specific parameters here if needed
73        }
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}