Skip to main content

roder_core/
teams.rs

1use std::collections::HashMap;
2use std::path::PathBuf;
3
4use roder_api::events::{ThreadId, TurnId};
5use roder_api::policy_mode::PolicyMode;
6use roder_api::teams::{
7    AgentTeamDisplayMode, TeamId, TeamMailboxMessage, TeamMailboxMessageKind, TeamMemberDescriptor,
8    TeamMemberId, TeamMemberRole, TeamMemberStatus, TeamTaskDescriptor,
9};
10use serde::{Deserialize, Serialize};
11use time::OffsetDateTime;
12use tokio::sync::RwLock;
13
14#[derive(Debug, Clone)]
15pub struct TeamStartRequest {
16    pub lead_thread_id: Option<ThreadId>,
17    pub display_mode: AgentTeamDisplayMode,
18    pub members: Vec<TeamMemberStartRequest>,
19}
20
21#[derive(Debug, Clone)]
22pub struct TeamMemberStartRequest {
23    pub name: String,
24    pub model_provider: Option<String>,
25    pub model: Option<String>,
26}
27
28#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
29#[serde(rename_all = "camelCase")]
30pub struct TeamState {
31    pub id: TeamId,
32    pub lead_thread_id: ThreadId,
33    pub display_mode: AgentTeamDisplayMode,
34    pub members: Vec<TeamMemberDescriptor>,
35    pub mailbox: Vec<TeamMailboxMessage>,
36    pub tasks: Vec<TeamTaskDescriptor>,
37    #[serde(with = "time::serde::rfc3339")]
38    pub created_at: OffsetDateTime,
39    #[serde(with = "time::serde::rfc3339")]
40    pub updated_at: OffsetDateTime,
41}
42
43#[derive(Debug)]
44pub struct TeamManager {
45    teams: RwLock<HashMap<TeamId, TeamState>>,
46    mailbox_reservations: RwLock<HashMap<String, TurnId>>,
47    data_dir: PathBuf,
48}
49
50impl Default for TeamManager {
51    fn default() -> Self {
52        Self::new(default_team_data_dir())
53    }
54}
55
56impl TeamManager {
57    pub fn new(data_dir: PathBuf) -> Self {
58        Self {
59            teams: RwLock::new(HashMap::new()),
60            mailbox_reservations: RwLock::new(HashMap::new()),
61            data_dir,
62        }
63    }
64
65    pub async fn insert(&self, team: TeamState) -> anyhow::Result<TeamState> {
66        let mut teams = self.teams.write().await;
67        self.persist(&team).await?;
68        teams.insert(team.id.clone(), team.clone());
69        Ok(team)
70    }
71
72    pub async fn get(&self, team_id: &str) -> Option<TeamState> {
73        if let Some(team) = self.teams.read().await.get(team_id).cloned() {
74            return Some(team);
75        }
76        let mut teams = self.teams.write().await;
77        self.load_locked(&mut teams, team_id).await.ok().flatten()
78    }
79
80    pub async fn list(&self) -> Vec<TeamState> {
81        let mut teams = self.teams.write().await;
82        let _ = self.load_all_locked(&mut teams).await;
83        let mut listed = teams.values().cloned().collect::<Vec<_>>();
84        listed.sort_by_key(|team| std::cmp::Reverse(team.updated_at));
85        listed
86    }
87
88    pub async fn update_member(
89        &self,
90        team_id: &str,
91        member_id: &str,
92        update: impl FnOnce(&mut TeamMemberDescriptor),
93    ) -> anyhow::Result<TeamState> {
94        let mut teams = self.teams.write().await;
95        let mut team = self
96            .load_locked(&mut teams, team_id)
97            .await?
98            .ok_or_else(|| anyhow::anyhow!("unknown team {team_id:?}"))?;
99        let member = team
100            .members
101            .iter_mut()
102            .find(|member| member.id == member_id)
103            .ok_or_else(|| anyhow::anyhow!("unknown team member {member_id:?}"))?;
104        update(member);
105        team.updated_at = OffsetDateTime::now_utc();
106        self.persist(&team).await?;
107        teams.insert(team.id.clone(), team.clone());
108        Ok(team)
109    }
110
111    pub async fn add_member(
112        &self,
113        team_id: &str,
114        member: TeamMemberDescriptor,
115    ) -> anyhow::Result<TeamState> {
116        let mut teams = self.teams.write().await;
117        let mut team = self
118            .load_locked(&mut teams, team_id)
119            .await?
120            .ok_or_else(|| anyhow::anyhow!("unknown team {team_id:?}"))?;
121        anyhow::ensure!(
122            !team.members.iter().any(|existing| existing.id == member.id),
123            "team member id {:?} already exists",
124            member.id
125        );
126        if let Some(agent_path) = member.agent_path.as_deref() {
127            anyhow::ensure!(
128                !team
129                    .members
130                    .iter()
131                    .any(|existing| existing.agent_path.as_deref() == Some(agent_path)),
132                "agent path {agent_path:?} already exists"
133            );
134        }
135        team.members.push(member);
136        team.updated_at = OffsetDateTime::now_utc();
137        self.persist(&team).await?;
138        teams.insert(team.id.clone(), team.clone());
139        Ok(team)
140    }
141
142    pub async fn append_mailbox_message(
143        &self,
144        team_id: &str,
145        from_member_id: Option<TeamMemberId>,
146        to_member_id: TeamMemberId,
147        kind: TeamMailboxMessageKind,
148        text: String,
149    ) -> anyhow::Result<TeamState> {
150        let mut teams = self.teams.write().await;
151        let mut team = self
152            .load_locked(&mut teams, team_id)
153            .await?
154            .ok_or_else(|| anyhow::anyhow!("unknown team {team_id:?}"))?;
155        if !team.members.iter().any(|member| member.id == to_member_id) {
156            anyhow::bail!("unknown team member {to_member_id:?}");
157        }
158        team.mailbox.push(TeamMailboxMessage {
159            id: uuid::Uuid::new_v4().to_string(),
160            team_id: team_id.to_string(),
161            from_member_id,
162            to_member_id,
163            kind,
164            text,
165            delivered: false,
166            timestamp: OffsetDateTime::now_utc(),
167        });
168        team.updated_at = OffsetDateTime::now_utc();
169        self.persist(&team).await?;
170        teams.insert(team.id.clone(), team.clone());
171        Ok(team)
172    }
173
174    /// Reserve currently pending messages for one active turn without marking
175    /// them delivered. Reservations are process-local: a crash makes the
176    /// messages eligible for redelivery, preserving at-least-once delivery.
177    pub async fn reserve_pending_mailbox_messages(
178        &self,
179        team_id: &str,
180        member_id: &str,
181        turn_id: &TurnId,
182    ) -> anyhow::Result<Vec<TeamMailboxMessage>> {
183        let mut teams = self.teams.write().await;
184        let team = self
185            .load_locked(&mut teams, team_id)
186            .await?
187            .ok_or_else(|| anyhow::anyhow!("unknown team {team_id:?}"))?;
188        anyhow::ensure!(
189            team.members.iter().any(|member| member.id == member_id),
190            "unknown team member {member_id:?}"
191        );
192        let mut reservations = self.mailbox_reservations.write().await;
193        let pending = team
194            .mailbox
195            .into_iter()
196            .filter(|message| {
197                message.to_member_id == member_id
198                    && !message.delivered
199                    && !reservations.contains_key(&message.id)
200            })
201            .collect::<Vec<_>>();
202        for message in &pending {
203            reservations.insert(message.id.clone(), turn_id.clone());
204        }
205        Ok(pending)
206    }
207
208    /// Whether any undelivered mailbox message is addressed to this member.
209    ///
210    /// Read-only on purpose: unlike [`Self::reserve_pending_mailbox_messages`]
211    /// it takes no reservation, so a waiter can ask "is there anything for me?"
212    /// without consuming the queue it is about to be woken for.
213    pub async fn has_pending_mailbox_messages(&self, team_id: &str, member_id: &str) -> bool {
214        let mut teams = self.teams.write().await;
215        let Ok(Some(team)) = self.load_locked(&mut teams, team_id).await else {
216            return false;
217        };
218        team.mailbox
219            .into_iter()
220            .any(|message| message.to_member_id == member_id && !message.delivered)
221    }
222
223    pub async fn release_mailbox_reservations_for_turn(&self, turn_id: &TurnId) {
224        self.mailbox_reservations
225            .write()
226            .await
227            .retain(|_, reserved_turn_id| reserved_turn_id != turn_id);
228    }
229
230    pub async fn release_mailbox_reservations(&self, turn_id: &TurnId, message_ids: &[String]) {
231        let message_ids = message_ids.iter().collect::<std::collections::HashSet<_>>();
232        self.mailbox_reservations
233            .write()
234            .await
235            .retain(|message_id, reserved_turn_id| {
236                reserved_turn_id != turn_id || !message_ids.contains(message_id)
237            });
238    }
239
240    pub async fn mark_mailbox_messages_delivered(
241        &self,
242        team_id: &str,
243        turn_id: &TurnId,
244        message_ids: &[String],
245    ) -> anyhow::Result<()> {
246        if message_ids.is_empty() {
247            return Ok(());
248        }
249        let mut teams = self.teams.write().await;
250        let mut team = self
251            .load_locked(&mut teams, team_id)
252            .await?
253            .ok_or_else(|| anyhow::anyhow!("unknown team {team_id:?}"))?;
254        let mut reservations = self.mailbox_reservations.write().await;
255        let message_ids = message_ids
256            .iter()
257            .filter(|message_id| reservations.get(*message_id) == Some(turn_id))
258            .collect::<std::collections::HashSet<_>>();
259        let mut changed = false;
260        for message in &mut team.mailbox {
261            if message_ids.contains(&message.id) && !message.delivered {
262                message.delivered = true;
263                changed = true;
264            }
265        }
266        if changed {
267            team.updated_at = OffsetDateTime::now_utc();
268            self.persist(&team).await?;
269            teams.insert(team.id.clone(), team);
270        }
271        reservations.retain(|message_id, _| !message_ids.contains(message_id));
272        Ok(())
273    }
274
275    pub async fn set_member_policy_mode(
276        &self,
277        team_id: &str,
278        member_id: &str,
279        policy_mode: PolicyMode,
280    ) -> anyhow::Result<TeamState> {
281        self.update_member(team_id, member_id, |member| {
282            member.policy_mode = policy_mode;
283        })
284        .await
285    }
286
287    pub async fn policy_mode_for_thread(&self, thread_id: &str) -> Option<PolicyMode> {
288        self.list()
289            .await
290            .into_iter()
291            .flat_map(|team| team.members.into_iter())
292            .find(|member| member.thread_id == thread_id)
293            .map(|member| member.policy_mode)
294    }
295
296    pub async fn member_for_thread(
297        &self,
298        thread_id: &str,
299    ) -> Option<(TeamId, TeamMemberDescriptor)> {
300        self.list().await.into_iter().find_map(|team| {
301            let team_id = team.id;
302            team.members
303                .into_iter()
304                .find(|member| member.thread_id == thread_id)
305                .map(|member| (team_id, member))
306        })
307    }
308
309    pub async fn complete_member_turn(
310        &self,
311        thread_id: &str,
312        turn_id: &str,
313        status: TeamMemberStatus,
314        final_message: Option<String>,
315        terminal_error: Option<String>,
316    ) -> anyhow::Result<Option<(TeamId, TeamMemberDescriptor)>> {
317        let mut teams = self.teams.write().await;
318        self.load_all_locked(&mut teams).await?;
319        let Some(team_id) = teams.iter().find_map(|(team_id, team)| {
320            team.members
321                .iter()
322                .any(|member| {
323                    member.thread_id == thread_id
324                        && member.current_turn_id.as_deref() == Some(turn_id)
325                })
326                .then(|| team_id.clone())
327        }) else {
328            return Ok(None);
329        };
330        let mut team = teams
331            .get(&team_id)
332            .cloned()
333            .expect("team found by member turn");
334        let member = team
335            .members
336            .iter_mut()
337            .find(|member| {
338                member.thread_id == thread_id && member.current_turn_id.as_deref() == Some(turn_id)
339            })
340            .expect("member located by team_for_member_turn");
341        member.status = status;
342        member.current_turn_id = None;
343        member.final_message = final_message;
344        member.terminal_error = terminal_error;
345        let completed = member.clone();
346        team.updated_at = OffsetDateTime::now_utc();
347        self.persist(&team).await?;
348        teams.insert(team_id.clone(), team);
349        Ok(Some((team_id, completed)))
350    }
351
352    pub async fn remove(&self, team_id: &str) -> anyhow::Result<Option<TeamState>> {
353        let mut teams = self.teams.write().await;
354        let _ = self.load_locked(&mut teams, team_id).await?;
355        let path = self.team_file(team_id);
356        if tokio::fs::try_exists(&path).await.unwrap_or(false) {
357            tokio::fs::remove_file(path).await?;
358        }
359        Ok(teams.remove(team_id))
360    }
361
362    async fn persist(&self, team: &TeamState) -> anyhow::Result<()> {
363        tokio::fs::create_dir_all(&self.data_dir).await?;
364        let data = serde_json::to_vec_pretty(team)?;
365        let temp_path = self
366            .data_dir
367            .join(format!(".team-{}.tmp", uuid::Uuid::new_v4()));
368        tokio::fs::write(&temp_path, data).await?;
369        if let Err(err) = tokio::fs::rename(&temp_path, self.team_file(&team.id)).await {
370            let _ = tokio::fs::remove_file(&temp_path).await;
371            return Err(err.into());
372        }
373        Ok(())
374    }
375
376    async fn load_locked(
377        &self,
378        teams: &mut HashMap<TeamId, TeamState>,
379        team_id: &str,
380    ) -> anyhow::Result<Option<TeamState>> {
381        if let Some(team) = teams.get(team_id).cloned() {
382            return Ok(Some(team));
383        }
384        let path = self.team_file(team_id);
385        if !tokio::fs::try_exists(&path).await.unwrap_or(false) {
386            return Ok(None);
387        }
388        let data = tokio::fs::read(path).await?;
389        let team = serde_json::from_slice::<TeamState>(&data)?;
390        anyhow::ensure!(
391            team.id == team_id,
392            "persisted team id {:?} does not match file id {team_id:?}",
393            team.id
394        );
395        teams.insert(team.id.clone(), team.clone());
396        Ok(Some(team))
397    }
398
399    async fn load_all_locked(&self, teams: &mut HashMap<TeamId, TeamState>) -> anyhow::Result<()> {
400        if !tokio::fs::try_exists(&self.data_dir).await.unwrap_or(false) {
401            return Ok(());
402        }
403        let mut entries = tokio::fs::read_dir(&self.data_dir).await?;
404        while let Some(entry) = entries.next_entry().await? {
405            let path = entry.path();
406            if path.extension().and_then(|extension| extension.to_str()) != Some("json") {
407                continue;
408            }
409            let Ok(data) = tokio::fs::read(&path).await else {
410                continue;
411            };
412            let Ok(team) = serde_json::from_slice::<TeamState>(&data) else {
413                continue;
414            };
415            if teams.contains_key(&team.id) {
416                continue;
417            }
418            teams.insert(team.id.clone(), team);
419        }
420        Ok(())
421    }
422
423    fn team_file(&self, team_id: &str) -> PathBuf {
424        self.data_dir.join(format!("{team_id}.json"))
425    }
426}
427
428pub(crate) fn lead_member(
429    thread_id: ThreadId,
430    model_provider: Option<String>,
431    model: Option<String>,
432    policy_mode: PolicyMode,
433) -> TeamMemberDescriptor {
434    TeamMemberDescriptor {
435        id: "lead".to_string(),
436        role: TeamMemberRole::Lead,
437        name: "Lead".to_string(),
438        task_name: Some("root".to_string()),
439        agent_path: Some("/root".to_string()),
440        thread_id,
441        parent_thread_id: None,
442        current_turn_id: None,
443        model_provider,
444        model,
445        policy_mode,
446        status: TeamMemberStatus::Idle,
447        final_message: None,
448        terminal_error: None,
449        pane_id: None,
450    }
451}
452
453pub(crate) fn teammate_member(
454    id: TeamMemberId,
455    name: String,
456    thread_id: ThreadId,
457    model_provider: Option<String>,
458    model: Option<String>,
459    policy_mode: PolicyMode,
460) -> TeamMemberDescriptor {
461    TeamMemberDescriptor {
462        id,
463        role: TeamMemberRole::Teammate,
464        name,
465        task_name: None,
466        agent_path: None,
467        thread_id,
468        parent_thread_id: None,
469        current_turn_id: None,
470        model_provider,
471        model,
472        policy_mode,
473        status: TeamMemberStatus::Idle,
474        final_message: None,
475        terminal_error: None,
476        pane_id: None,
477    }
478}
479
480pub(crate) fn default_team_data_dir() -> PathBuf {
481    std::env::var_os("RODER_DATA_DIR")
482        .map(PathBuf::from)
483        .or_else(|| std::env::var_os("HOME").map(|home| PathBuf::from(home).join(".roder")))
484        .unwrap_or_else(|| PathBuf::from(".roder"))
485        .join("teams")
486}