khora_core/agent/
completion.rs1use std::collections::HashMap;
20use std::sync::OnceLock;
21
22use crate::control::gorna::AgentId;
23use crate::renderer::api::core::StageHandle;
24
25pub struct AgentDone;
27
28#[derive(Debug, Clone, Copy, PartialEq, Eq)]
30pub enum CompletionOutcome {
31 Completed,
33 Skipped,
35}
36
37struct AgentCompletion {
40 handle: StageHandle<AgentDone>,
41 outcome: OnceLock<CompletionOutcome>,
42}
43
44impl AgentCompletion {
45 fn new() -> Self {
46 Self {
47 handle: StageHandle::<AgentDone>::default(),
48 outcome: OnceLock::new(),
49 }
50 }
51}
52
53pub struct AgentCompletionMap {
59 entries: HashMap<AgentId, AgentCompletion>,
60}
61
62impl AgentCompletionMap {
63 pub fn new(agent_ids: &[AgentId]) -> Self {
65 let mut entries = HashMap::with_capacity(agent_ids.len());
66 for id in agent_ids {
67 entries.insert(*id, AgentCompletion::new());
68 }
69 Self { entries }
70 }
71
72 pub fn mark(&self, id: AgentId, outcome: CompletionOutcome) -> bool {
77 match self.entries.get(&id) {
78 Some(entry) => {
79 let _ = entry.outcome.set(outcome);
80 entry.handle.mark_done();
81 true
82 }
83 None => false,
84 }
85 }
86
87 pub fn outcome(&self, id: AgentId) -> Option<CompletionOutcome> {
90 self.entries.get(&id)?.outcome.get().copied()
91 }
92
93 pub fn is_done(&self, id: AgentId) -> bool {
95 self.entries
96 .get(&id)
97 .is_some_and(|entry| entry.handle.is_done())
98 }
99
100 pub async fn wait(&self, id: AgentId) -> Option<CompletionOutcome> {
104 let entry = self.entries.get(&id)?;
105 entry.handle.wait().await;
106 entry.outcome.get().copied()
107 }
108
109 pub fn known_ids(&self) -> impl Iterator<Item = AgentId> + '_ {
111 self.entries.keys().copied()
112 }
113}
114
115#[cfg(test)]
116mod tests {
117 use super::*;
118
119 fn rt() -> tokio::runtime::Runtime {
120 tokio::runtime::Builder::new_current_thread()
121 .enable_all()
122 .build()
123 .unwrap()
124 }
125
126 #[test]
127 fn wait_returns_immediately_after_mark() {
128 let map = AgentCompletionMap::new(&[AgentId::Renderer]);
129 map.mark(AgentId::Renderer, CompletionOutcome::Completed);
130
131 let result = rt().block_on(async { map.wait(AgentId::Renderer).await });
132 assert_eq!(result, Some(CompletionOutcome::Completed));
133 }
134
135 #[test]
136 fn wait_unblocks_when_mark_arrives_later() {
137 use std::sync::Arc;
138
139 let map = Arc::new(AgentCompletionMap::new(&[AgentId::ShadowRenderer]));
140 let map_clone = Arc::clone(&map);
141
142 let result = rt().block_on(async move {
143 let waiter = tokio::spawn(async move { map_clone.wait(AgentId::ShadowRenderer).await });
144 tokio::task::yield_now().await;
146 map.mark(AgentId::ShadowRenderer, CompletionOutcome::Skipped);
147 waiter.await.unwrap()
148 });
149 assert_eq!(result, Some(CompletionOutcome::Skipped));
150 }
151
152 #[test]
153 fn wait_on_unknown_agent_returns_none() {
154 let map = AgentCompletionMap::new(&[AgentId::Renderer]);
155 let result = rt().block_on(async { map.wait(AgentId::Audio).await });
156 assert_eq!(result, None);
157 }
158
159 #[test]
160 fn outcome_distinguishes_completed_from_skipped() {
161 let map = AgentCompletionMap::new(&[AgentId::Renderer, AgentId::ShadowRenderer]);
162 map.mark(AgentId::Renderer, CompletionOutcome::Completed);
163 map.mark(AgentId::ShadowRenderer, CompletionOutcome::Skipped);
164
165 assert_eq!(
166 map.outcome(AgentId::Renderer),
167 Some(CompletionOutcome::Completed)
168 );
169 assert_eq!(
170 map.outcome(AgentId::ShadowRenderer),
171 Some(CompletionOutcome::Skipped)
172 );
173 assert_eq!(map.outcome(AgentId::Physics), None);
174 }
175
176 #[test]
177 fn mark_is_idempotent() {
178 let map = AgentCompletionMap::new(&[AgentId::Renderer]);
179 assert!(map.mark(AgentId::Renderer, CompletionOutcome::Completed));
180 assert!(map.mark(AgentId::Renderer, CompletionOutcome::Skipped));
183 assert_eq!(
184 map.outcome(AgentId::Renderer),
185 Some(CompletionOutcome::Completed)
186 );
187 }
188
189 #[test]
190 fn mark_unknown_agent_returns_false() {
191 let map = AgentCompletionMap::new(&[AgentId::Renderer]);
192 assert!(!map.mark(AgentId::Audio, CompletionOutcome::Completed));
193 }
194}