use crate::client::OpenAIClient;
use crate::config::OpenAIConfig;
use crate::error::OpenAIAgentError;
use crate::models::{ChatMessage, ChatRequest, ChatResponse, ToolCall};
use crate::tools::ToolRegistry;
use crate::websocket_client::{WebSocketClient, RealtimeEvent, ServerEvent};
use std::sync::Arc;
#[derive(Debug, Clone)]
pub struct AgentState {
pub(crate) messages: Vec<ChatMessage>,
pub(crate) token_count: usize,
}
impl AgentState {
pub fn token_count(&self) -> usize {
self.token_count
}
pub fn message_count(&self) -> usize {
self.messages.len()
}
pub fn messages(&self) -> impl Iterator<Item = &ChatMessage> {
self.messages.iter()
}
}
pub struct Agent {
client: OpenAIClient,
tools: Arc<ToolRegistry>,
state: AgentState,
max_turns: usize,
websocket_client: Option<WebSocketClient>,
}
impl Agent {
#[doc(hidden)]
pub(crate) fn from_builder(builder: AgentBuilder) -> Result<Self, OpenAIAgentError> {
let client = OpenAIClient::new(builder.config.clone().unwrap_or_default())?;
let state = AgentState {
messages: builder.messages,
token_count: 0,
};
Ok(Self {
client,
tools: builder.tools,
state,
max_turns: builder.max_turns,
websocket_client: builder.websocket_client,
})
}
pub async fn run(&mut self, input: impl Into<String>) -> Result<String, OpenAIAgentError> {
self.state.messages.push(ChatMessage::user(input.into()));
let mut turns = 0;
let mut final_response = String::new();
while turns < self.max_turns {
turns += 1;
let request = self.prepare_request()?;
let response = self.client.chat_completion(request).await?;
if let Some(usage) = response.usage.as_ref() {
self.state.token_count += usage.total_tokens;
}
if let Some(choice) = response.choices.first() {
self.state.messages.push(choice.message.clone());
if let Some(tool_calls) = &choice.message.tool_calls {
if !tool_calls.is_empty() {
for tool_call in tool_calls {
let result_msg = self.execute_tool_call(tool_call).await?;
self.state.messages.push(result_msg);
}
continue;
}
}
if let Some(content) = &choice.message.content {
if !content.trim().is_empty() {
final_response = content.clone();
return Ok(final_response);
}
}
if choice.finish_reason == "tool_calls" {
continue;
}
return Err(OpenAIAgentError::Parse(
format!("Assistant returned empty message with finish_reason: {}", choice.finish_reason),
));
} else {
return Err(OpenAIAgentError::Parse(
"No response choices received".to_string(),
));
}
}
Err(OpenAIAgentError::Agent(format!(
"Agent exceeded maximum turns ({})",
self.max_turns
)))
}
async fn execute_tool_call(&self, tc: &ToolCall) -> Result<ChatMessage, OpenAIAgentError> {
let tool_name = &tc.function.name;
let arguments = &tc.function.arguments;
let tool_call_id = &tc.id;
let tool = self
.tools
.get(tool_name)
.ok_or_else(|| OpenAIAgentError::Tool(format!("Tool not found: {}", tool_name)))?;
let parsed_args = serde_json::from_str(arguments)
.map_err(|e| OpenAIAgentError::Parse(format!("Failed to parse tool arguments: {}", e)))?;
let result = tool.execute(parsed_args).await?;
let response = ChatMessage {
role: "tool".to_string(),
content: Some(result),
name: Some(tool_name.clone()),
tool_call_id: Some(tool_call_id.clone()),
tool_calls: None,
};
Ok(response)
}
fn prepare_request(&self) -> Result<ChatRequest, OpenAIAgentError> {
let config = self.client.config();
let mut request = ChatRequest {
model: config.model().to_string(),
messages: self.state.messages.clone(),
tools: None,
max_tokens: Some(config.max_tokens()),
temperature: Some(config.temperature()),
response_format: None,
stream: Some(config.stream()),
};
if !self.tools.is_empty() {
request.tools = Some(self.tools.definitions());
}
Ok(request)
}
pub fn state(&self) -> &AgentState {
&self.state
}
pub fn push_user_message(&mut self, content: impl Into<String>) {
self.state.messages.push(ChatMessage::user(content.into()));
}
pub fn push_assistant_message(&mut self, content: impl Into<String>) {
self.state.messages.push(ChatMessage::assistant(content.into()));
}
pub async fn connect_realtime(&mut self, model_name: &str) -> Result<(), OpenAIAgentError> {
let ws_client = self
.websocket_client
.as_mut()
.ok_or_else(|| OpenAIAgentError::Agent(
"No WebSocket client configured (call `with_websocket()` first).".to_string()
))?;
ws_client.connect(model_name).await
}
pub async fn send_realtime_event(
&mut self,
event: &RealtimeEvent
) -> Result<(), OpenAIAgentError> {
let ws_client = self
.websocket_client
.as_mut()
.ok_or_else(|| OpenAIAgentError::Agent(
"No WebSocket client configured (call `with_websocket()` first).".to_string()
))?;
ws_client.send_event(event).await
}
pub async fn process_realtime_events<F>(
&mut self,
on_event: F
) -> Result<(), OpenAIAgentError>
where
F: FnMut(ServerEvent) -> Result<(), OpenAIAgentError>
{
let ws_client = self
.websocket_client
.as_mut()
.ok_or_else(|| OpenAIAgentError::Agent(
"No WebSocket client configured (call `with_websocket()` first).".to_string()
))?;
ws_client.process_incoming(on_event).await
}
pub async fn close_realtime(&mut self) -> Result<(), OpenAIAgentError> {
if let Some(ws_client) = &mut self.websocket_client {
ws_client.close().await?;
}
Ok(())
}
}
pub struct AgentBuilder {
pub(crate) config: Option<OpenAIConfig>,
pub(crate) tools: Arc<ToolRegistry>,
pub(crate) messages: Vec<ChatMessage>,
pub(crate) max_turns: usize,
pub(crate) websocket_client: Option<WebSocketClient>,
}
impl AgentBuilder {
pub fn new() -> Self {
Self {
config: None,
tools: Arc::new(ToolRegistry::new()),
messages: Vec::new(),
max_turns: 10,
websocket_client: None,
}
}
pub fn with_config(mut self, config: OpenAIConfig) -> Self {
self.config = Some(config);
self
}
pub fn with_tools(mut self, tools: ToolRegistry) -> Self {
self.tools = Arc::new(tools);
self
}
pub fn with_system_prompt(mut self, prompt: impl Into<String>) -> Self {
self.messages.push(ChatMessage::system(prompt.into()));
self
}
pub fn with_message(mut self, message: ChatMessage) -> Self {
self.messages.push(message);
self
}
pub fn with_max_turns(mut self, max_turns: usize) -> Self {
self.max_turns = max_turns;
self
}
pub fn with_websocket(mut self) -> Result<Self, OpenAIAgentError> {
let cfg = self.config.clone().unwrap_or_default();
let ws_client = WebSocketClient::new(cfg)?;
self.websocket_client = Some(ws_client);
Ok(self)
}
pub fn build(self) -> Result<Agent, OpenAIAgentError> {
Agent::from_builder(self)
}
}
impl Default for AgentBuilder {
fn default() -> Self {
Self::new()
}
}