agent_base/engine/
session.rs1use 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 messages: Vec<Message>,
15 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}