use std::collections::{HashMap, hash_map as path_std_collections_hash_map};
use std::hash::{BuildHasher, Hasher};
use std::path::PathBuf;
use std::sync::atomic::{AtomicU64, Ordering};
use std::sync::{Arc, Mutex};
#[derive(Clone)]
enum WorkdirValue {
Valid(PathBuf),
Invalid,
ReplayFailed,
}
#[derive(Clone)]
pub(crate) enum WorkdirSnapshot {
Valid(PathBuf),
Invalid,
ReplayFailed,
}
#[derive(Clone)]
struct PendingWorkdirResult {
mutation_id: tau_proto::AgentMetadataMutationId,
expected_cwd: PathBuf,
identity: crate::tool_started_identity::ToolStartedIdentity,
lock_wait_duration_seconds: Option<u64>,
awaiting_echo: bool,
cancel_requested: bool,
}
#[derive(Clone)]
pub(crate) struct CompletedPendingWorkdir {
pub(crate) identity: crate::tool_started_identity::ToolStartedIdentity,
pub(crate) lock_wait_duration_seconds: Option<u64>,
pub(crate) matched_request: bool,
pub(crate) cancel_requested: bool,
}
#[derive(Clone)]
pub(crate) struct CwdState {
instance_name: Arc<Mutex<tau_proto::ExtensionName>>,
context_label: Arc<Mutex<String>>,
process_startup_cwd: Arc<Mutex<Result<PathBuf, String>>>,
workdir_by_agent: Arc<Mutex<HashMap<tau_proto::AgentId, WorkdirValue>>>,
pending_ready_by_agent: Arc<
Mutex<
HashMap<tau_proto::AgentId, (tau_proto::SessionId, tau_proto::AgentInitializationId)>,
>,
>,
initialization_by_agent: Arc<
Mutex<
HashMap<tau_proto::AgentId, (tau_proto::SessionId, tau_proto::AgentInitializationId)>,
>,
>,
pending_workdir_by_agent: Arc<Mutex<HashMap<tau_proto::AgentId, PendingWorkdirResult>>>,
next_mutation_id: Arc<AtomicU64>,
mutation_id_salt: u64,
}
impl CwdState {
pub(crate) fn new() -> Self {
Self::with_startup_cwd(Self::read_process_startup_cwd())
}
#[cfg(any(test, feature = "echo-agent"))]
pub(crate) fn new_with_startup_cwd(cwd: PathBuf) -> Self {
Self::with_startup_cwd(Self::validate_startup_cwd(cwd))
}
fn with_startup_cwd(process_startup_cwd: Result<PathBuf, String>) -> Self {
Self {
instance_name: Arc::new(Mutex::new(
tau_proto::ExtensionName::parse("core-shell")
.expect("known-safe default extension name must be valid"),
)),
context_label: Arc::new(Mutex::new("default".to_owned())),
process_startup_cwd: Arc::new(Mutex::new(process_startup_cwd)),
workdir_by_agent: Arc::new(Mutex::new(HashMap::new())),
pending_ready_by_agent: Arc::new(Mutex::new(HashMap::new())),
initialization_by_agent: Arc::new(Mutex::new(HashMap::new())),
pending_workdir_by_agent: Arc::new(Mutex::new(HashMap::new())),
next_mutation_id: Arc::new(AtomicU64::new(1)),
mutation_id_salt: {
let mut hasher = path_std_collections_hash_map::RandomState::new().build_hasher();
hasher.write_u64(0);
hasher.finish()
},
}
}
pub(crate) fn set_instance_name(&self, name: tau_proto::ExtensionName) {
*self
.instance_name
.lock()
.expect("cwd instance lock poisoned") = name;
}
pub(crate) fn set_context_label(&self, prefix: Option<&tau_proto::ToolNamePrefix>) {
*self
.context_label
.lock()
.expect("workdir context label lock poisoned") =
prefix.map_or_else(|| "default".to_owned(), |prefix| prefix.as_str().to_owned());
}
pub(crate) fn context_label(&self) -> String {
self.context_label
.lock()
.expect("workdir context label lock poisoned")
.clone()
}
pub(crate) fn key(&self) -> tau_proto::AgentMetadataKey {
let name = self
.instance_name
.lock()
.expect("cwd instance lock poisoned");
tau_proto::AgentMetadataKey::new(format!("ext_{}_cwd", name.as_str()))
}
pub(crate) fn get(&self, agent_id: &tau_proto::AgentId) -> Option<PathBuf> {
self.workdir_by_agent
.lock()
.expect("workdir map lock poisoned")
.get(agent_id)
.and_then(|value| match value {
WorkdirValue::Valid(path) => Some(path.clone()),
WorkdirValue::Invalid => None,
WorkdirValue::ReplayFailed => None,
})
}
fn read_process_startup_cwd() -> Result<PathBuf, String> {
std::env::current_dir()
.map_err(|error| format!("failed to read ext-shell process working directory: {error}"))
.and_then(Self::validate_startup_cwd)
}
fn validate_startup_cwd(cwd: PathBuf) -> Result<PathBuf, String> {
let cwd = cwd.canonicalize().map_err(|error| {
format!(
"failed to canonicalize ext-shell process working directory {}: {error}",
cwd.display()
)
})?;
if !cwd.is_dir() {
return Err(format!(
"ext-shell process working directory is not a directory: {}",
cwd.display()
));
}
Ok(cwd)
}
pub(crate) fn freeze_process_startup_cwd(&self) -> Result<(), String> {
let cwd = Self::read_process_startup_cwd();
*self
.process_startup_cwd
.lock()
.expect("process startup cwd lock poisoned") = cwd.clone();
cwd.map(|_| ())
}
pub(crate) fn process_default(&self) -> Result<PathBuf, String> {
self.process_startup_cwd
.lock()
.expect("process startup cwd lock poisoned")
.clone()
}
pub(crate) fn get_or_default(&self, agent_id: &tau_proto::AgentId) -> Result<PathBuf, String> {
match self.snapshot(agent_id)? {
WorkdirSnapshot::Valid(path) => Ok(path),
WorkdirSnapshot::Invalid => Err(
"remembered workdir metadata is invalid; repair it with an absolute workdir path"
.to_owned(),
),
WorkdirSnapshot::ReplayFailed => Err(
"workdir replay failed for this agent; reload the agent before retrying".to_owned(),
),
}
}
pub(crate) fn snapshot(
&self,
agent_id: &tau_proto::AgentId,
) -> Result<WorkdirSnapshot, String> {
match self
.workdir_by_agent
.lock()
.expect("workdir map lock poisoned")
.get(agent_id)
.cloned()
{
Some(WorkdirValue::Valid(path)) => Ok(WorkdirSnapshot::Valid(path)),
Some(WorkdirValue::Invalid) => Ok(WorkdirSnapshot::Invalid),
Some(WorkdirValue::ReplayFailed) => Ok(WorkdirSnapshot::ReplayFailed),
None => self.process_default().map(WorkdirSnapshot::Valid),
}
}
pub(crate) fn set(&self, agent_id: tau_proto::AgentId, cwd: PathBuf) {
self.workdir_by_agent
.lock()
.expect("workdir map lock poisoned")
.insert(agent_id, WorkdirValue::Valid(cwd));
}
pub(crate) fn set_metadata_text(&self, agent_id: tau_proto::AgentId, cwd: PathBuf) -> bool {
if cwd.as_os_str().is_empty() || !cwd.is_absolute() {
self.set_invalid(agent_id);
return false;
}
self.set(agent_id, cwd);
true
}
pub(crate) fn unset(&self, agent_id: &tau_proto::AgentId) {
self.workdir_by_agent
.lock()
.expect("workdir map lock poisoned")
.remove(agent_id);
}
pub(crate) fn set_invalid(&self, agent_id: tau_proto::AgentId) {
self.workdir_by_agent
.lock()
.expect("workdir map lock poisoned")
.insert(agent_id, WorkdirValue::Invalid);
}
pub(crate) fn set_replay_failed(&self, agent_id: tau_proto::AgentId) {
self.workdir_by_agent
.lock()
.expect("workdir map lock poisoned")
.insert(agent_id, WorkdirValue::ReplayFailed);
}
pub(crate) fn is_invalid(&self, agent_id: &tau_proto::AgentId) -> bool {
matches!(
self.workdir_by_agent
.lock()
.expect("workdir map lock poisoned")
.get(agent_id),
Some(WorkdirValue::Invalid)
)
}
pub(crate) fn is_replay_failed(&self, agent_id: &tau_proto::AgentId) -> bool {
matches!(
self.workdir_by_agent
.lock()
.expect("workdir map lock poisoned")
.get(agent_id),
Some(WorkdirValue::ReplayFailed)
)
}
pub(crate) fn set_pending_ready(
&self,
agent_id: tau_proto::AgentId,
session_id: tau_proto::SessionId,
agent_initialization_id: tau_proto::AgentInitializationId,
) {
self.initialization_by_agent
.lock()
.expect("cwd initialization map lock poisoned")
.insert(
agent_id.clone(),
(session_id.clone(), agent_initialization_id.clone()),
);
self.pending_ready_by_agent
.lock()
.expect("cwd ready map lock poisoned")
.insert(agent_id, (session_id, agent_initialization_id));
}
pub(crate) fn initialization(
&self,
agent_id: &tau_proto::AgentId,
) -> Option<(tau_proto::SessionId, tau_proto::AgentInitializationId)> {
self.initialization_by_agent
.lock()
.expect("cwd initialization map lock poisoned")
.get(agent_id)
.cloned()
}
pub(crate) fn remove_initialization(&self, agent_id: &tau_proto::AgentId) {
self.initialization_by_agent
.lock()
.expect("cwd initialization map lock poisoned")
.remove(agent_id);
}
pub(crate) fn take_pending_ready(
&self,
agent_id: &tau_proto::AgentId,
) -> Option<(tau_proto::SessionId, tau_proto::AgentInitializationId)> {
self.pending_ready_by_agent
.lock()
.expect("cwd ready map lock poisoned")
.remove(agent_id)
}
pub(crate) fn pending_ready(
&self,
agent_id: &tau_proto::AgentId,
) -> Option<(tau_proto::SessionId, tau_proto::AgentInitializationId)> {
self.pending_ready_by_agent
.lock()
.expect("cwd ready map lock poisoned")
.get(agent_id)
.cloned()
}
pub(crate) fn start_pending_workdir_result(
&self,
agent_id: tau_proto::AgentId,
expected_cwd: PathBuf,
identity: impl Into<crate::tool_started_identity::ToolStartedIdentity>,
lock_wait_duration_seconds: Option<u64>,
) -> Result<(), Box<crate::tool_started_identity::ToolStartedIdentity>> {
let identity = identity.into();
let mut pending = self
.pending_workdir_by_agent
.lock()
.expect("workdir setter map lock poisoned");
if pending.contains_key(&agent_id) {
return Err(Box::new(identity));
}
pending.insert(
agent_id,
PendingWorkdirResult {
mutation_id: {
let mut hasher =
path_std_collections_hash_map::RandomState::new().build_hasher();
hasher.write_u64(self.mutation_id_salt);
hasher.write_u64(self.next_mutation_id.fetch_add(1, Ordering::Relaxed));
tau_proto::AgentMetadataMutationId::parse(format!(
"ext-shell-workdir-{:016x}",
hasher.finish()
))
.expect("fixed-size mutation id is valid")
},
expected_cwd,
identity,
lock_wait_duration_seconds,
awaiting_echo: false,
cancel_requested: false,
},
);
Ok(())
}
pub(crate) fn pending_workdir_mutation_id(
&self,
agent_id: &tau_proto::AgentId,
call_id: &tau_proto::ToolCallId,
) -> Option<tau_proto::AgentMetadataMutationId> {
self.pending_workdir_by_agent
.lock()
.expect("workdir setter map lock poisoned")
.get(agent_id)
.filter(|pending| &pending.identity.call_id == call_id)
.map(|pending| pending.mutation_id.clone())
}
pub(crate) fn mark_pending_workdir_awaiting_echo(
&self,
agent_id: &tau_proto::AgentId,
call_id: &tau_proto::ToolCallId,
) -> bool {
let mut pending = self
.pending_workdir_by_agent
.lock()
.expect("workdir setter map lock poisoned");
let Some(pending) = pending.get_mut(agent_id) else {
return false;
};
if &pending.identity.call_id != call_id {
return false;
}
pending.awaiting_echo = true;
true
}
pub(crate) fn pending_workdir_target(
&self,
agent_id: &tau_proto::AgentId,
call_id: &tau_proto::ToolCallId,
) -> Option<PathBuf> {
self.pending_workdir_by_agent
.lock()
.expect("workdir setter map lock poisoned")
.get(agent_id)
.filter(|pending| &pending.identity.call_id == call_id)
.map(|pending| pending.expected_cwd.clone())
}
pub(crate) fn committed_pending_workdir_result(
&self,
agent_id: &tau_proto::AgentId,
committed_cwd: &PathBuf,
mutation_id: Option<&tau_proto::AgentMetadataMutationId>,
) -> Option<CompletedPendingWorkdir> {
let pending_by_agent = self
.pending_workdir_by_agent
.lock()
.expect("workdir setter map lock poisoned");
let pending = pending_by_agent.get(agent_id)?;
if !pending.awaiting_echo || mutation_id != Some(&pending.mutation_id) {
return None;
}
Some(CompletedPendingWorkdir {
matched_request: pending.expected_cwd == *committed_cwd,
cancel_requested: pending.cancel_requested,
identity: pending.identity.clone(),
lock_wait_duration_seconds: pending.lock_wait_duration_seconds,
})
}
pub(crate) fn correlated_pending_workdir_result(
&self,
agent_id: &tau_proto::AgentId,
mutation_id: Option<&tau_proto::AgentMetadataMutationId>,
) -> Option<CompletedPendingWorkdir> {
let pending_by_agent = self
.pending_workdir_by_agent
.lock()
.expect("workdir setter map lock poisoned");
let pending = pending_by_agent.get(agent_id)?;
if !pending.awaiting_echo || mutation_id != Some(&pending.mutation_id) {
return None;
}
Some(CompletedPendingWorkdir {
matched_request: false,
cancel_requested: pending.cancel_requested,
identity: pending.identity.clone(),
lock_wait_duration_seconds: pending.lock_wait_duration_seconds,
})
}
pub(crate) fn take_pending_workdir_result(
&self,
agent_id: &tau_proto::AgentId,
) -> Option<CompletedPendingWorkdir> {
let pending = self
.pending_workdir_by_agent
.lock()
.expect("workdir setter map lock poisoned")
.remove(agent_id)?;
Some(CompletedPendingWorkdir {
matched_request: false,
cancel_requested: pending.cancel_requested,
identity: pending.identity,
lock_wait_duration_seconds: pending.lock_wait_duration_seconds,
})
}
pub(crate) fn take_pending_workdir_by_call(
&self,
call_id: &tau_proto::ToolCallId,
) -> Option<CompletedPendingWorkdir> {
let mut pending = self
.pending_workdir_by_agent
.lock()
.expect("workdir setter map lock poisoned");
let agent_id = pending.iter().find_map(|(agent_id, item)| {
(&item.identity.call_id == call_id).then(|| agent_id.clone())
})?;
let pending = pending.remove(&agent_id)?;
Some(CompletedPendingWorkdir {
matched_request: false,
cancel_requested: pending.cancel_requested,
identity: pending.identity,
lock_wait_duration_seconds: pending.lock_wait_duration_seconds,
})
}
pub(crate) fn take_all_pending_workdirs(&self) -> Vec<CompletedPendingWorkdir> {
self.pending_workdir_by_agent
.lock()
.expect("workdir setter map lock poisoned")
.drain()
.map(|(_, pending)| CompletedPendingWorkdir {
matched_request: false,
cancel_requested: pending.cancel_requested,
identity: pending.identity,
lock_wait_duration_seconds: pending.lock_wait_duration_seconds,
})
.collect()
}
pub(crate) fn request_pending_workdir_cancel(&self, call_id: &tau_proto::ToolCallId) -> bool {
let mut pending = self
.pending_workdir_by_agent
.lock()
.expect("workdir setter map lock poisoned");
let Some(item) = pending
.values_mut()
.find(|item| &item.identity.call_id == call_id && item.awaiting_echo)
else {
return false;
};
item.cancel_requested = true;
true
}
}