Skip to main content

agent_base/engine/
session.rs

1use std::collections::HashSet;
2
3use serde::{Deserialize, Serialize};
4
5use crate::types::{ChatMessage, ImageAttachment, Message, MessageRole, ToolCallMessage};
6
7use crate::types::SessionId;
8
9#[derive(Clone, Debug, Default, Serialize, Deserialize)]
10pub struct AgentSession {
11    id: Option<SessionId>,
12    /// Simplified message history for internal logic and debugging.
13    /// Each entry corresponds to a `ChatMessage` in `chat_messages`.
14    messages: Vec<Message>,
15    /// LLM API format messages, sent directly to the provider.
16    /// This is the source of truth for the conversation state.
17    chat_messages: Vec<ChatMessage>,
18    always_allowed_actions: HashSet<String>,
19}
20
21impl AgentSession {
22    pub fn new(id: SessionId) -> Self {
23        Self {
24            id: Some(id),
25            messages: Vec::new(),
26            chat_messages: Vec::new(),
27            always_allowed_actions: HashSet::new(),
28        }
29    }
30
31    pub fn id(&self) -> Option<SessionId> {
32        self.id.clone()
33    }
34
35    pub fn messages(&self) -> &[Message] {
36        &self.messages
37    }
38
39    pub fn chat_messages(&self) -> &[ChatMessage] {
40        &self.chat_messages
41    }
42
43    pub fn is_action_allowed(&self, action_key: &str) -> bool {
44        self.always_allowed_actions.contains(action_key)
45    }
46
47    pub fn allow_action(&mut self, action_key: impl Into<String>) {
48        self.always_allowed_actions.insert(action_key.into());
49    }
50
51    pub fn push_message(&mut self, role: MessageRole, content: impl Into<String>) {
52        let content = content.into();
53        self.messages.push(Message {
54            role: role.clone(),
55            content: content.clone(),
56        });
57        let chat_msg = match role {
58            MessageRole::System => ChatMessage::system(content),
59            MessageRole::User => ChatMessage::user(content),
60            MessageRole::Assistant => ChatMessage::assistant(content),
61            MessageRole::Tool => ChatMessage::tool(String::new(), content),
62        };
63        self.chat_messages.push(chat_msg);
64    }
65
66    pub fn push_user_message_with_images(
67        &mut self,
68        content: impl Into<String>,
69        images: Vec<ImageAttachment>,
70    ) {
71        let content = content.into();
72        self.messages.push(Message {
73            role: MessageRole::User,
74            content: content.clone(),
75        });
76        self.chat_messages
77            .push(ChatMessage::user_with_images(content, images));
78    }
79
80    pub fn push_assistant_tool_call(
81        &mut self,
82        tool_call_id: &str,
83        tool_name: &str,
84        arguments_json: &str,
85    ) {
86        self.chat_messages.push(ChatMessage::assistant_tool_call(
87            tool_call_id,
88            tool_name,
89            arguments_json,
90        ));
91    }
92
93    pub fn push_assistant_tool_calls(&mut self, tool_calls: &[(String, String, String)]) {
94        let calls: Vec<ToolCallMessage> = tool_calls
95            .iter()
96            .map(|(id, name, args)| ToolCallMessage {
97                id: id.clone(),
98                name: name.clone(),
99                arguments: args.clone(),
100            })
101            .collect();
102        self.chat_messages.push(ChatMessage::Assistant {
103            content: None,
104            reasoning_content: None,
105            tool_calls: Some(calls),
106        });
107    }
108
109    pub fn push_tool_result(&mut self, tool_call_id: &str, content: impl Into<String>) {
110        let content = content.into();
111        self.messages.push(Message {
112            role: MessageRole::Tool,
113            content: content.clone(),
114        });
115        self.chat_messages.push(ChatMessage::tool(
116            tool_call_id,
117            content,
118        ));
119    }
120
121    pub fn close_dangling_tool_calls(&mut self, error_summary: &str) {
122        let assistant_idx = self.chat_messages.iter().rposition(|m| {
123            matches!(m, ChatMessage::Assistant { tool_calls: Some(tc), .. } if !tc.is_empty())
124        });
125
126        let Some(assistant_idx) = assistant_idx else {
127            return;
128        };
129
130        let ChatMessage::Assistant { tool_calls: Some(tc), .. } = &self.chat_messages[assistant_idx] else {
131            return;
132        };
133
134        let all_ids: Vec<String> = tc.iter().map(|t| t.id.clone()).collect();
135
136        let answered_ids: Vec<String> = self.chat_messages[assistant_idx + 1..]
137            .iter()
138            .filter_map(|m| match m {
139                ChatMessage::Tool { tool_call_id, .. } => Some(tool_call_id.clone()),
140                _ => None,
141            })
142            .collect();
143
144        for id in &all_ids {
145            if !answered_ids.iter().any(|a| a == id) {
146                self.push_tool_result(id, error_summary);
147            }
148        }
149    }
150}