use std::collections::HashSet;
use tau_proto::{AgentPromptId, ContextItem, ToolCallId};
#[derive(Clone, Default)]
pub(crate) struct AgentActivity {
optimistic_submissions: usize,
active_prompts: HashSet<AgentPromptId>,
active_tools: HashSet<ToolCallId>,
backgrounded_tools: HashSet<ToolCallId>,
}
impl AgentActivity {
pub(crate) fn is_in_progress(&self) -> bool {
self.optimistic_submissions != 0
|| !self.active_prompts.is_empty()
|| !self.active_tools.is_empty()
}
pub(crate) fn has_active_prompts(&self) -> bool {
!self.active_prompts.is_empty()
}
pub(crate) fn mark_optimistic_submission(&mut self) {
self.optimistic_submissions = self.optimistic_submissions.saturating_add(1);
}
pub(crate) fn start_prompt(&mut self, agent_prompt_id: &AgentPromptId) {
if self.active_prompts.contains(agent_prompt_id) {
return;
}
self.optimistic_submissions = self.optimistic_submissions.saturating_sub(1);
self.active_prompts.insert(agent_prompt_id.clone());
}
pub(crate) fn finish_prompt(
&mut self,
agent_prompt_id: &AgentPromptId,
output_items: &[ContextItem],
) {
if self.finish_active_prompt(
agent_prompt_id,
tool_call_ids_from_output_items(output_items),
) {
return;
}
self.optimistic_submissions = self.optimistic_submissions.saturating_sub(1);
}
pub(crate) fn finish_prompt_with_tool_call_ids<'a>(
&mut self,
agent_prompt_id: &AgentPromptId,
call_ids: impl IntoIterator<Item = &'a ToolCallId>,
) {
if self.finish_active_prompt(agent_prompt_id, call_ids) {
return;
}
self.optimistic_submissions = self.optimistic_submissions.saturating_sub(1);
}
#[cfg(test)]
pub(crate) fn finish_prompt_if_active(
&mut self,
agent_prompt_id: &AgentPromptId,
output_items: &[ContextItem],
) {
let _ = self.finish_active_prompt(
agent_prompt_id,
tool_call_ids_from_output_items(output_items),
);
}
pub(crate) fn finish_prompt_if_active_with_tool_call_ids<'a>(
&mut self,
agent_prompt_id: &AgentPromptId,
call_ids: impl IntoIterator<Item = &'a ToolCallId>,
) {
let _ = self.finish_active_prompt(agent_prompt_id, call_ids);
}
fn finish_active_prompt<'a>(
&mut self,
agent_prompt_id: &AgentPromptId,
call_ids: impl IntoIterator<Item = &'a ToolCallId>,
) -> bool {
if self.active_prompts.remove(agent_prompt_id) {
for call_id in call_ids {
self.active_tools.insert(call_id.clone());
}
true
} else {
false
}
}
pub(crate) fn start_tool(&mut self, call_id: &ToolCallId) {
self.active_tools.insert(call_id.clone());
}
pub(crate) fn background_tool(&mut self, call_id: &ToolCallId) {
self.backgrounded_tools.insert(call_id.clone());
self.active_tools.insert(call_id.clone());
}
pub(crate) fn finish_tool(&mut self, call_id: &ToolCallId) {
if !self.backgrounded_tools.contains(call_id) {
self.active_tools.remove(call_id);
}
}
pub(crate) fn finish_background_tool(&mut self, call_id: &ToolCallId) {
self.backgrounded_tools.remove(call_id);
self.active_tools.remove(call_id);
}
pub(crate) fn clear_optimistic_submissions(&mut self) {
self.optimistic_submissions = 0;
}
pub(crate) fn clear(&mut self) {
self.optimistic_submissions = 0;
self.active_prompts.clear();
self.active_tools.clear();
self.backgrounded_tools.clear();
}
}
fn tool_call_ids_from_output_items(
output_items: &[ContextItem],
) -> impl Iterator<Item = &ToolCallId> {
output_items.iter().filter_map(|item| match item {
ContextItem::ToolCall(call) => Some(&call.call_id),
_ => None,
})
}