Skip to main content

roder_core/
goals.rs

1use std::collections::HashMap;
2use std::path::PathBuf;
3use std::sync::Arc;
4
5use anyhow::Context;
6use roder_api::events::{RoderEvent, ThreadId};
7use roder_api::goals::{
8    ThreadGoal, ThreadGoalCleared, ThreadGoalController, ThreadGoalPatch, ThreadGoalStatus,
9    ThreadGoalUpdated, validate_thread_goal_budget, validate_thread_goal_objective,
10};
11use roder_api::inference::InstructionBundle;
12use roder_api::thread::ThreadStore;
13use time::{Duration, OffsetDateTime};
14
15mod accounting;
16mod prompts;
17mod runtime;
18#[cfg(test)]
19mod tests;
20
21pub(crate) use prompts::continuation_prompt;
22use prompts::objective_updated_prompt;
23use tokio::sync::Mutex;
24
25use crate::bus::EventBus;
26use crate::runtime::{Runtime, StartTurnRequest};
27
28const GOAL_STATE_FILE: &str = "goal.json";
29
30#[derive(Debug, Default)]
31struct GoalCache {
32    goals: HashMap<ThreadId, Option<ThreadGoal>>,
33    turns: HashMap<String, accounting::GoalTurnProgress>,
34    empty_turns: HashMap<ThreadId, u8>,
35    continuation_requests: HashMap<ThreadId, StartTurnRequest>,
36}
37
38#[derive(Clone)]
39pub struct RuntimeGoalController {
40    bus: EventBus,
41    thread_store: Option<Arc<dyn ThreadStore>>,
42    thread_root: Option<PathBuf>,
43    cache: Arc<Mutex<GoalCache>>,
44    mutation: Arc<Mutex<()>>,
45}
46
47impl RuntimeGoalController {
48    pub fn new(bus: EventBus, thread_store: Option<Arc<dyn ThreadStore>>) -> Self {
49        let thread_root = thread_store
50            .as_ref()
51            .and_then(|store| store.local_thread_root());
52        Self {
53            bus,
54            thread_store,
55            thread_root,
56            cache: Arc::new(Mutex::new(GoalCache::default())),
57            mutation: Arc::new(Mutex::new(())),
58        }
59    }
60
61    pub async fn apply_goal_instructions(
62        &self,
63        thread_id: &ThreadId,
64        mut instructions: InstructionBundle,
65        mode: roder_api::policy_mode::PolicyMode,
66    ) -> anyhow::Result<InstructionBundle> {
67        let Some(goal) = self.get_thread_goal(thread_id).await? else {
68            return Ok(instructions);
69        };
70        if !matches!(
71            goal.status,
72            ThreadGoalStatus::Active | ThreadGoalStatus::BudgetLimited
73        ) {
74            return Ok(instructions);
75        }
76        let mut addition = if goal.status == ThreadGoalStatus::BudgetLimited {
77            if !self
78                .cache
79                .lock()
80                .await
81                .turns
82                .values()
83                .any(|progress| progress.matches_goal(&goal))
84            {
85                return Ok(instructions);
86            }
87            prompts::budget_limit_prompt(&goal)
88        } else {
89            continuation_prompt(&goal)
90        };
91        addition.push_str("\n\n");
92        addition.push_str(prompts::permission_prompt(mode));
93        instructions.developer = Some(match instructions.developer {
94            Some(existing) if !existing.trim().is_empty() => format!("{existing}\n\n{addition}"),
95            _ => addition,
96        });
97        Ok(instructions)
98    }
99
100    /// Account work explicitly attributed to the active goal (including compaction).
101    pub async fn account_turn_usage(
102        &self,
103        thread_id: &ThreadId,
104        tokens_used: i64,
105        elapsed: Duration,
106    ) -> anyhow::Result<Option<ThreadGoal>> {
107        let _guard = self.mutation.lock().await;
108        let Some(mut goal) = self.load_goal(thread_id).await? else {
109            return Ok(None);
110        };
111        if goal.status != ThreadGoalStatus::Active {
112            return Ok(Some(goal));
113        }
114        goal.tokens_used = goal.tokens_used.saturating_add(tokens_used.max(0));
115        goal.time_used_seconds = goal
116            .time_used_seconds
117            .saturating_add(elapsed.whole_seconds().max(0));
118        enforce_budget(&mut goal);
119        goal.updated_at = OffsetDateTime::now_utc();
120        self.store_goal(goal.clone()).await?;
121        self.emit_goal_updated(goal.clone()).await;
122        Ok(Some(goal))
123    }
124
125    pub async fn active_goal(&self, thread_id: &ThreadId) -> anyhow::Result<Option<ThreadGoal>> {
126        Ok(self
127            .get_thread_goal(thread_id)
128            .await?
129            .filter(|goal| goal.status == ThreadGoalStatus::Active))
130    }
131
132    async fn load_goal(&self, thread_id: &ThreadId) -> anyhow::Result<Option<ThreadGoal>> {
133        let mut cache = self.cache.lock().await;
134        if let Some(goal) = cache.goals.get(thread_id) {
135            return Ok(goal.clone());
136        }
137        let goal = match self.goal_path(thread_id) {
138            Some(path) if path.exists() => {
139                let bytes = tokio::fs::read(&path)
140                    .await
141                    .with_context(|| format!("read goal state {}", path.display()))?;
142                Some(
143                    serde_json::from_slice::<ThreadGoal>(&bytes)
144                        .with_context(|| format!("parse goal state {}", path.display()))?,
145                )
146            }
147            _ => None,
148        };
149        cache.goals.insert(thread_id.clone(), goal.clone());
150        Ok(goal)
151    }
152
153    async fn store_goal(&self, goal: ThreadGoal) -> anyhow::Result<()> {
154        if let Some(path) = self.goal_path(&goal.thread_id) {
155            if let Some(parent) = path.parent() {
156                tokio::fs::create_dir_all(parent)
157                    .await
158                    .with_context(|| format!("create goal directory {}", parent.display()))?;
159            }
160            let bytes = serde_json::to_vec_pretty(&goal).context("serialize goal state")?;
161            let temporary = path.with_extension(format!("{}.tmp", uuid::Uuid::new_v4()));
162            tokio::fs::write(&temporary, bytes)
163                .await
164                .with_context(|| format!("write goal state {}", temporary.display()))?;
165            tokio::fs::rename(&temporary, &path)
166                .await
167                .with_context(|| format!("commit goal state {}", path.display()))?;
168        }
169        self.cache
170            .lock()
171            .await
172            .goals
173            .insert(goal.thread_id.clone(), Some(goal));
174        Ok(())
175    }
176
177    async fn remove_goal(&self, thread_id: &ThreadId) -> anyhow::Result<bool> {
178        let existed = self.load_goal(thread_id).await?.is_some();
179        if let Some(path) = self.goal_path(thread_id)
180            && path.exists()
181        {
182            tokio::fs::remove_file(&path)
183                .await
184                .with_context(|| format!("remove goal state {}", path.display()))?;
185        }
186        self.cache
187            .lock()
188            .await
189            .goals
190            .insert(thread_id.clone(), None);
191        Ok(existed)
192    }
193
194    fn goal_path(&self, thread_id: &ThreadId) -> Option<PathBuf> {
195        self.thread_root
196            .as_ref()
197            .map(|root| root.join(thread_id).join(GOAL_STATE_FILE))
198    }
199
200    pub(crate) async fn emit_goal_updated(&self, goal: ThreadGoal) {
201        let event = RoderEvent::ThreadGoalUpdated(ThreadGoalUpdated {
202            thread_id: goal.thread_id.clone(),
203            goal,
204            timestamp: OffsetDateTime::now_utc(),
205        });
206        self.emit_goal_event(event).await;
207    }
208
209    async fn emit_goal_cleared(&self, thread_id: ThreadId) {
210        let event = RoderEvent::ThreadGoalCleared(ThreadGoalCleared {
211            thread_id,
212            timestamp: OffsetDateTime::now_utc(),
213        });
214        self.emit_goal_event(event).await;
215    }
216
217    async fn emit_goal_event(&self, event: RoderEvent) {
218        let envelope = self.bus.emit(event);
219        if let (Some(store), Some(thread_id)) = (&self.thread_store, envelope.thread_id.as_ref()) {
220            let _ = store.append_event(thread_id, &envelope).await;
221        }
222    }
223}
224
225#[async_trait::async_trait]
226impl ThreadGoalController for RuntimeGoalController {
227    async fn get_thread_goal(&self, thread_id: &ThreadId) -> anyhow::Result<Option<ThreadGoal>> {
228        let _guard = self.mutation.lock().await;
229        self.flush_thread_progress(thread_id).await?;
230        self.load_goal(thread_id).await
231    }
232
233    async fn create_thread_goal(
234        &self,
235        thread_id: &ThreadId,
236        objective: String,
237        token_budget: Option<i64>,
238    ) -> anyhow::Result<ThreadGoal> {
239        let objective = objective.trim().to_string();
240        validate_thread_goal_objective(&objective)?;
241        validate_thread_goal_budget(token_budget)?;
242        let _guard = self.mutation.lock().await;
243        if self
244            .load_goal(thread_id)
245            .await?
246            .is_some_and(|goal| goal.status != ThreadGoalStatus::Complete)
247        {
248            anyhow::bail!(
249                "cannot create a new goal because this thread has an unfinished goal; complete the existing goal first"
250            );
251        }
252        let now = OffsetDateTime::now_utc();
253        let goal = ThreadGoal {
254            thread_id: thread_id.clone(),
255            objective,
256            status: ThreadGoalStatus::Active,
257            token_budget,
258            tokens_used: 0,
259            time_used_seconds: 0,
260            created_at: now,
261            updated_at: now,
262        };
263        self.store_goal(goal.clone()).await?;
264        self.attach_goal_to_turns(&goal).await;
265        self.emit_goal_updated(goal.clone()).await;
266        Ok(goal)
267    }
268
269    async fn set_thread_goal(
270        &self,
271        thread_id: &ThreadId,
272        patch: ThreadGoalPatch,
273    ) -> anyhow::Result<Option<ThreadGoal>> {
274        if let Some(objective) = patch.objective.as_deref() {
275            validate_thread_goal_objective(objective)?;
276        }
277        if let Some(token_budget) = patch.token_budget {
278            validate_thread_goal_budget(token_budget)?;
279        }
280        let _guard = self.mutation.lock().await;
281        self.flush_thread_progress(thread_id).await?;
282        let now = OffsetDateTime::now_utc();
283        let mut goal = match self.load_goal(thread_id).await? {
284            Some(goal) => goal,
285            None => {
286                let Some(objective) = patch.objective.as_ref() else {
287                    return Ok(None);
288                };
289                ThreadGoal {
290                    thread_id: thread_id.clone(),
291                    objective: objective.trim().to_string(),
292                    status: ThreadGoalStatus::Active,
293                    token_budget: None,
294                    tokens_used: 0,
295                    time_used_seconds: 0,
296                    created_at: now,
297                    updated_at: now,
298                }
299            }
300        };
301        if let Some(objective) = patch.objective {
302            goal.objective = objective.trim().to_string();
303        }
304        if let Some(status) = patch.status {
305            // Exhausted budgets take precedence over a requested pause or blocked state.
306            if goal.status != ThreadGoalStatus::BudgetLimited
307                || !matches!(status, ThreadGoalStatus::Paused | ThreadGoalStatus::Blocked)
308            {
309                goal.status = status;
310            }
311        }
312        if let Some(token_budget) = patch.token_budget {
313            goal.token_budget = token_budget;
314        }
315        enforce_budget(&mut goal);
316        goal.updated_at = OffsetDateTime::now_utc();
317        self.store_goal(goal.clone()).await?;
318        self.attach_goal_to_turns(&goal).await;
319        self.emit_goal_updated(goal.clone()).await;
320        Ok(Some(goal))
321    }
322
323    async fn clear_thread_goal(&self, thread_id: &ThreadId) -> anyhow::Result<bool> {
324        let _guard = self.mutation.lock().await;
325        let cleared = self.remove_goal(thread_id).await?;
326        self.detach_thread_goal(thread_id).await;
327        if cleared {
328            self.emit_goal_cleared(thread_id.clone()).await;
329        }
330        Ok(cleared)
331    }
332}
333
334fn enforce_budget(goal: &mut ThreadGoal) {
335    if goal.status == ThreadGoalStatus::Active
336        && goal
337            .token_budget
338            .is_some_and(|budget| goal.tokens_used >= budget)
339    {
340        goal.status = ThreadGoalStatus::BudgetLimited;
341    }
342}
343
344pub(crate) fn status_after_error(error: &anyhow::Error) -> ThreadGoalStatus {
345    use roder_api::provider_error::{ProviderFailure, ProviderFailureKind};
346    if error
347        .downcast_ref::<ProviderFailure>()
348        .is_some_and(|failure| {
349            matches!(
350                failure.kind,
351                ProviderFailureKind::UsageLimit | ProviderFailureKind::QuotaExceeded
352            )
353        })
354    {
355        ThreadGoalStatus::UsageLimited
356    } else {
357        ThreadGoalStatus::Blocked
358    }
359}