Skip to main content

roder_core/
teams.rs

1use std::collections::HashMap;
2use std::path::PathBuf;
3
4use roder_api::events::ThreadId;
5use roder_api::policy_mode::PolicyMode;
6use roder_api::teams::{
7    AgentTeamDisplayMode, TeamId, TeamMailboxMessage, TeamMemberDescriptor, TeamMemberId,
8    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    data_dir: PathBuf,
47}
48
49impl Default for TeamManager {
50    fn default() -> Self {
51        Self::new(default_team_data_dir())
52    }
53}
54
55impl TeamManager {
56    pub fn new(data_dir: PathBuf) -> Self {
57        Self {
58            teams: RwLock::new(HashMap::new()),
59            data_dir,
60        }
61    }
62
63    pub async fn insert(&self, team: TeamState) -> anyhow::Result<TeamState> {
64        self.persist(&team).await?;
65        self.teams
66            .write()
67            .await
68            .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        self.load(team_id).await.ok().flatten()
77    }
78
79    pub async fn list(&self) -> Vec<TeamState> {
80        let mut teams = self
81            .teams
82            .read()
83            .await
84            .values()
85            .cloned()
86            .collect::<Vec<_>>();
87        teams.sort_by_key(|team| std::cmp::Reverse(team.updated_at));
88        teams
89    }
90
91    pub async fn update_member(
92        &self,
93        team_id: &str,
94        member_id: &str,
95        update: impl FnOnce(&mut TeamMemberDescriptor),
96    ) -> anyhow::Result<TeamState> {
97        let mut team = self
98            .get(team_id)
99            .await
100            .ok_or_else(|| anyhow::anyhow!("unknown team {team_id:?}"))?;
101        let member = team
102            .members
103            .iter_mut()
104            .find(|member| member.id == member_id)
105            .ok_or_else(|| anyhow::anyhow!("unknown team member {member_id:?}"))?;
106        update(member);
107        team.updated_at = OffsetDateTime::now_utc();
108        self.insert(team).await
109    }
110
111    pub async fn append_mailbox_message(
112        &self,
113        team_id: &str,
114        from_member_id: Option<TeamMemberId>,
115        to_member_id: TeamMemberId,
116        text: String,
117    ) -> anyhow::Result<TeamState> {
118        let mut team = self
119            .get(team_id)
120            .await
121            .ok_or_else(|| anyhow::anyhow!("unknown team {team_id:?}"))?;
122        if !team.members.iter().any(|member| member.id == to_member_id) {
123            anyhow::bail!("unknown team member {to_member_id:?}");
124        }
125        team.mailbox.push(TeamMailboxMessage {
126            id: uuid::Uuid::new_v4().to_string(),
127            team_id: team_id.to_string(),
128            from_member_id,
129            to_member_id,
130            text,
131            timestamp: OffsetDateTime::now_utc(),
132        });
133        team.updated_at = OffsetDateTime::now_utc();
134        self.insert(team).await
135    }
136
137    pub async fn set_member_policy_mode(
138        &self,
139        team_id: &str,
140        member_id: &str,
141        policy_mode: PolicyMode,
142    ) -> anyhow::Result<TeamState> {
143        self.update_member(team_id, member_id, |member| {
144            member.policy_mode = policy_mode;
145        })
146        .await
147    }
148
149    pub async fn policy_mode_for_thread(&self, thread_id: &str) -> Option<PolicyMode> {
150        self.teams
151            .read()
152            .await
153            .values()
154            .flat_map(|team| team.members.iter())
155            .find(|member| member.thread_id == thread_id)
156            .map(|member| member.policy_mode)
157    }
158
159    pub async fn member_for_thread(
160        &self,
161        thread_id: &str,
162    ) -> Option<(TeamId, TeamMemberDescriptor)> {
163        self.teams.read().await.values().find_map(|team| {
164            team.members
165                .iter()
166                .find(|member| member.thread_id == thread_id)
167                .cloned()
168                .map(|member| (team.id.clone(), member))
169        })
170    }
171
172    pub async fn complete_member_turn(
173        &self,
174        thread_id: &str,
175        turn_id: &str,
176        status: TeamMemberStatus,
177    ) -> anyhow::Result<Option<(TeamId, TeamMemberDescriptor)>> {
178        let Some(mut team) = self.team_for_member_turn(thread_id, turn_id).await else {
179            return Ok(None);
180        };
181        let team_id = team.id.clone();
182        let member = team
183            .members
184            .iter_mut()
185            .find(|member| {
186                member.thread_id == thread_id && member.current_turn_id.as_deref() == Some(turn_id)
187            })
188            .expect("member located by team_for_member_turn");
189        member.status = status;
190        member.current_turn_id = None;
191        let completed = member.clone();
192        team.updated_at = OffsetDateTime::now_utc();
193        self.insert(team).await?;
194        Ok(Some((team_id, completed)))
195    }
196
197    pub async fn remove(&self, team_id: &str) -> anyhow::Result<Option<TeamState>> {
198        let removed = self.teams.write().await.remove(team_id);
199        let path = self.team_file(team_id);
200        if tokio::fs::try_exists(&path).await.unwrap_or(false) {
201            tokio::fs::remove_file(path).await?;
202        }
203        Ok(removed)
204    }
205
206    async fn persist(&self, team: &TeamState) -> anyhow::Result<()> {
207        tokio::fs::create_dir_all(&self.data_dir).await?;
208        let data = serde_json::to_vec_pretty(team)?;
209        tokio::fs::write(self.team_file(&team.id), data).await?;
210        Ok(())
211    }
212
213    async fn load(&self, team_id: &str) -> anyhow::Result<Option<TeamState>> {
214        let path = self.team_file(team_id);
215        if !tokio::fs::try_exists(&path).await.unwrap_or(false) {
216            return Ok(None);
217        }
218        let data = tokio::fs::read(path).await?;
219        let team = serde_json::from_slice::<TeamState>(&data)?;
220        self.teams
221            .write()
222            .await
223            .insert(team.id.clone(), team.clone());
224        Ok(Some(team))
225    }
226
227    async fn team_for_member_turn(&self, thread_id: &str, turn_id: &str) -> Option<TeamState> {
228        self.teams
229            .read()
230            .await
231            .values()
232            .find(|team| {
233                team.members.iter().any(|member| {
234                    member.thread_id == thread_id
235                        && member.current_turn_id.as_deref() == Some(turn_id)
236                })
237            })
238            .cloned()
239    }
240
241    fn team_file(&self, team_id: &str) -> PathBuf {
242        self.data_dir.join(format!("{team_id}.json"))
243    }
244}
245
246pub(crate) fn lead_member(
247    thread_id: ThreadId,
248    model_provider: Option<String>,
249    model: Option<String>,
250    policy_mode: PolicyMode,
251) -> TeamMemberDescriptor {
252    TeamMemberDescriptor {
253        id: "lead".to_string(),
254        role: TeamMemberRole::Lead,
255        name: "Lead".to_string(),
256        thread_id,
257        current_turn_id: None,
258        model_provider,
259        model,
260        policy_mode,
261        status: TeamMemberStatus::Idle,
262        pane_id: None,
263    }
264}
265
266pub(crate) fn teammate_member(
267    id: TeamMemberId,
268    name: String,
269    thread_id: ThreadId,
270    model_provider: Option<String>,
271    model: Option<String>,
272    policy_mode: PolicyMode,
273) -> TeamMemberDescriptor {
274    TeamMemberDescriptor {
275        id,
276        role: TeamMemberRole::Teammate,
277        name,
278        thread_id,
279        current_turn_id: None,
280        model_provider,
281        model,
282        policy_mode,
283        status: TeamMemberStatus::Idle,
284        pane_id: None,
285    }
286}
287
288pub(crate) fn default_team_data_dir() -> PathBuf {
289    std::env::var_os("RODER_DATA_DIR")
290        .map(PathBuf::from)
291        .or_else(|| std::env::var_os("HOME").map(|home| PathBuf::from(home).join(".roder")))
292        .unwrap_or_else(|| PathBuf::from(".roder"))
293        .join("teams")
294}