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 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 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}