1use async_trait::async_trait;
2use rlmesh_proto::common::v1::MessageBytes;
3
4use super::{
5 ActionReceivedEvent, EnvConnectedEvent, EpisodeCompletedEvent, EpisodeStartedEvent, HookError,
6 LogEvent, ModelConnectedEvent, ObservationEmittedEvent, RuntimeHooks, SessionEndedEvent,
7 SessionFailedEvent, SessionStartedEvent, StepCompletedEvent, TelemetrySummaryEvent,
8 TelemetryWindowEvent,
9};
10
11#[derive(Default)]
17pub struct RuntimeHookChain {
18 hooks: Vec<std::sync::Arc<dyn RuntimeHooks>>,
19}
20
21impl RuntimeHookChain {
22 pub fn new(hooks: Vec<std::sync::Arc<dyn RuntimeHooks>>) -> Self {
24 Self { hooks }
25 }
26
27 pub fn empty() -> Self {
29 Self::default()
30 }
31
32 pub fn len(&self) -> usize {
34 self.hooks.len()
35 }
36
37 pub fn is_empty(&self) -> bool {
39 self.hooks.is_empty()
40 }
41}
42
43#[async_trait]
44impl RuntimeHooks for RuntimeHookChain {
45 async fn env_connected(&self, event: EnvConnectedEvent) -> Result<(), HookError> {
46 let mut first_error = None;
47 for hook in &self.hooks {
48 if let Err(error) = hook.env_connected(event.clone()).await {
49 first_error.get_or_insert(error);
50 }
51 }
52 first_error.map_or(Ok(()), Err)
53 }
54
55 async fn model_connected(&self, event: ModelConnectedEvent) -> Result<(), HookError> {
56 let mut first_error = None;
57 for hook in &self.hooks {
58 if let Err(error) = hook.model_connected(event.clone()).await {
59 first_error.get_or_insert(error);
60 }
61 }
62 first_error.map_or(Ok(()), Err)
63 }
64
65 async fn session_started(&self, event: SessionStartedEvent) -> Result<(), HookError> {
66 let mut first_error = None;
67 for hook in &self.hooks {
68 if let Err(error) = hook.session_started(event.clone()).await {
69 first_error.get_or_insert(error);
70 }
71 }
72 first_error.map_or(Ok(()), Err)
73 }
74
75 async fn episode_started(&self, event: EpisodeStartedEvent) -> Result<(), HookError> {
76 let mut first_error = None;
77 for hook in &self.hooks {
78 if let Err(error) = hook.episode_started(event.clone()).await {
79 first_error.get_or_insert(error);
80 }
81 }
82 first_error.map_or(Ok(()), Err)
83 }
84
85 async fn episode_completed(&self, event: EpisodeCompletedEvent) -> Result<(), HookError> {
86 let mut first_error = None;
87 for hook in &self.hooks {
88 if let Err(error) = hook.episode_completed(event.clone()).await {
89 first_error.get_or_insert(error);
90 }
91 }
92 first_error.map_or(Ok(()), Err)
93 }
94
95 async fn action_received(&self, event: ActionReceivedEvent) -> Result<(), HookError> {
96 let mut first_error = None;
97 for hook in &self.hooks {
98 if let Err(error) = hook.action_received(event.clone()).await {
99 first_error.get_or_insert(error);
100 }
101 }
102 first_error.map_or(Ok(()), Err)
103 }
104
105 async fn transform_action(
106 &self,
107 event: ActionReceivedEvent,
108 ) -> Result<Option<MessageBytes>, HookError> {
109 let ActionReceivedEvent {
110 session_id,
111 route,
112 episode_id,
113 episode_record_id,
114 episode_ids,
115 episode_record_ids,
116 step,
117 env_index,
118 action_space,
119 mut action,
120 } = event;
121 for hook in &self.hooks {
122 action = hook
123 .transform_action(ActionReceivedEvent {
124 session_id: session_id.clone(),
125 route: route.clone(),
126 episode_id: episode_id.clone(),
127 episode_record_id: episode_record_id.clone(),
128 episode_ids: episode_ids.clone(),
129 episode_record_ids: episode_record_ids.clone(),
130 step,
131 env_index,
132 action_space: action_space.clone(),
133 action,
134 })
135 .await?;
136 }
137 Ok(action)
138 }
139
140 async fn step_completed(&self, event: StepCompletedEvent) -> Result<(), HookError> {
141 let mut first_error = None;
142 for hook in &self.hooks {
143 if let Err(error) = hook.step_completed(event.clone()).await {
144 first_error.get_or_insert(error);
145 }
146 }
147 first_error.map_or(Ok(()), Err)
148 }
149
150 async fn observation_emitted(&self, event: ObservationEmittedEvent) -> Result<(), HookError> {
151 let mut first_error = None;
152 for hook in &self.hooks {
153 if let Err(error) = hook.observation_emitted(event.clone()).await {
154 first_error.get_or_insert(error);
155 }
156 }
157 first_error.map_or(Ok(()), Err)
158 }
159
160 async fn transform_observation(
161 &self,
162 event: ObservationEmittedEvent,
163 ) -> Result<Option<MessageBytes>, HookError> {
164 let ObservationEmittedEvent {
165 session_id,
166 route,
167 episode_id,
168 episode_record_id,
169 episode_ids,
170 episode_record_ids,
171 step,
172 env_index,
173 is_reset,
174 num_envs,
175 observation_space,
176 mut observation,
177 } = event;
178 for hook in &self.hooks {
179 observation = hook
180 .transform_observation(ObservationEmittedEvent {
181 session_id: session_id.clone(),
182 route: route.clone(),
183 episode_id: episode_id.clone(),
184 episode_record_id: episode_record_id.clone(),
185 episode_ids: episode_ids.clone(),
186 episode_record_ids: episode_record_ids.clone(),
187 step,
188 env_index,
189 is_reset,
190 num_envs,
191 observation_space: observation_space.clone(),
192 observation,
193 })
194 .await?;
195 }
196 Ok(observation)
197 }
198
199 async fn telemetry_window(&self, event: TelemetryWindowEvent) -> Result<(), HookError> {
200 let mut first_error = None;
201 for hook in &self.hooks {
202 if let Err(error) = hook.telemetry_window(event.clone()).await {
203 first_error.get_or_insert(error);
204 }
205 }
206 first_error.map_or(Ok(()), Err)
207 }
208
209 async fn telemetry_summary(&self, event: TelemetrySummaryEvent) -> Result<(), HookError> {
210 let mut first_error = None;
211 for hook in &self.hooks {
212 if let Err(error) = hook.telemetry_summary(event.clone()).await {
213 first_error.get_or_insert(error);
214 }
215 }
216 first_error.map_or(Ok(()), Err)
217 }
218
219 async fn session_ended(&self, event: SessionEndedEvent) -> Result<(), HookError> {
220 let mut first_error = None;
221 for hook in &self.hooks {
222 if let Err(error) = hook.session_ended(event.clone()).await {
223 first_error.get_or_insert(error);
224 }
225 }
226 first_error.map_or(Ok(()), Err)
227 }
228
229 async fn session_failed(&self, event: SessionFailedEvent) -> Result<(), HookError> {
230 let mut first_error = None;
231 for hook in &self.hooks {
232 if let Err(error) = hook.session_failed(event.clone()).await {
233 first_error.get_or_insert(error);
234 }
235 }
236 first_error.map_or(Ok(()), Err)
237 }
238
239 async fn log(&self, event: LogEvent) -> Result<(), HookError> {
240 let mut first_error = None;
241 for hook in &self.hooks {
242 if let Err(error) = hook.log(event.clone()).await {
243 first_error.get_or_insert(error);
244 }
245 }
246 first_error.map_or(Ok(()), Err)
247 }
248}
249
250#[cfg(test)]
251mod tests {
252 use std::sync::{Arc, Mutex};
253
254 use async_trait::async_trait;
255 use rlmesh_proto::common::v1::MessageBytes;
256 use rlmesh_proto::spaces::v1::SpaceSpec;
257
258 use super::*;
259 use crate::hooks::{LogLevel, RuntimeRouteContext};
260
261 struct RecordingHook {
262 name: &'static str,
263 calls: Arc<Mutex<Vec<String>>>,
264 log_error: Option<&'static str>,
265 action_suffix: Option<u8>,
266 transform_error: Option<&'static str>,
267 }
268
269 impl RecordingHook {
270 fn new(name: &'static str, calls: Arc<Mutex<Vec<String>>>) -> Self {
271 Self {
272 name,
273 calls,
274 log_error: None,
275 action_suffix: None,
276 transform_error: None,
277 }
278 }
279
280 fn with_log_error(mut self, error: &'static str) -> Self {
281 self.log_error = Some(error);
282 self
283 }
284
285 fn with_action_suffix(mut self, suffix: u8) -> Self {
286 self.action_suffix = Some(suffix);
287 self
288 }
289
290 fn with_transform_error(mut self, error: &'static str) -> Self {
291 self.transform_error = Some(error);
292 self
293 }
294
295 fn record(&self, call: impl Into<String>) {
296 self.calls
297 .lock()
298 .expect("calls mutex poisoned")
299 .push(call.into());
300 }
301 }
302
303 #[async_trait]
304 impl RuntimeHooks for RecordingHook {
305 async fn log(&self, event: LogEvent) -> Result<(), HookError> {
306 self.record(format!("{}:log:{}", self.name, event.message));
307 if let Some(error) = self.log_error {
308 return Err(HookError::Message(error.to_string()));
309 }
310 Ok(())
311 }
312
313 async fn transform_action(
314 &self,
315 event: ActionReceivedEvent,
316 ) -> Result<Option<MessageBytes>, HookError> {
317 let data = event
318 .action
319 .as_ref()
320 .map(|action| action.data.clone())
321 .unwrap_or_default();
322 self.record(format!("{}:action:{data:?}", self.name));
323 if let Some(error) = self.transform_error {
324 return Err(HookError::Message(error.to_string()));
325 }
326 Ok(event.action.map(|mut action| {
327 if let Some(suffix) = self.action_suffix {
328 action.data.push(suffix);
329 }
330 action
331 }))
332 }
333 }
334
335 fn hook(hook: RecordingHook) -> Arc<dyn RuntimeHooks> {
336 Arc::new(hook)
337 }
338
339 fn recorded(calls: &Arc<Mutex<Vec<String>>>) -> Vec<String> {
340 calls.lock().expect("calls mutex poisoned").clone()
341 }
342
343 fn log_event() -> LogEvent {
344 LogEvent {
345 session_id: "session".to_string(),
346 route: RuntimeRouteContext::default(),
347 level: LogLevel::Info,
348 message: "hello".to_string(),
349 source: None,
350 }
351 }
352
353 fn action_event(data: Vec<u8>) -> ActionReceivedEvent {
354 ActionReceivedEvent {
355 session_id: "session".to_string(),
356 route: RuntimeRouteContext::default(),
357 episode_id: "episode".to_string(),
358 episode_record_id: "episode-artifact".to_string(),
359 episode_ids: vec!["episode".to_string()],
360 episode_record_ids: vec!["episode-artifact".to_string()],
361 step: 1,
362 env_index: 0,
363 action_space: SpaceSpec::default(),
364 action: Some(MessageBytes { data }),
365 }
366 }
367
368 #[tokio::test]
369 async fn event_hooks_call_every_hook_and_return_first_error() {
370 let calls = Arc::new(Mutex::new(Vec::new()));
371 let chain = RuntimeHookChain::new(vec![
372 hook(RecordingHook::new("first", calls.clone()).with_log_error("first failed")),
373 hook(RecordingHook::new("second", calls.clone()).with_log_error("second failed")),
374 hook(RecordingHook::new("third", calls.clone())),
375 ]);
376
377 let error = chain.log(log_event()).await.unwrap_err();
378
379 assert_eq!(error.to_string(), "first failed");
380 assert_eq!(
381 recorded(&calls),
382 vec!["first:log:hello", "second:log:hello", "third:log:hello"]
383 );
384 }
385
386 #[tokio::test]
387 async fn transform_hooks_run_in_order() {
388 let calls = Arc::new(Mutex::new(Vec::new()));
389 let chain = RuntimeHookChain::new(vec![
390 hook(RecordingHook::new("first", calls.clone()).with_action_suffix(1)),
391 hook(RecordingHook::new("second", calls.clone()).with_action_suffix(2)),
392 ]);
393
394 let action = chain
395 .transform_action(action_event(vec![0]))
396 .await
397 .unwrap()
398 .unwrap();
399
400 assert_eq!(action.data, vec![0, 1, 2]);
401 assert_eq!(
402 recorded(&calls),
403 vec!["first:action:[0]", "second:action:[0, 1]"]
404 );
405 }
406
407 #[tokio::test]
408 async fn transform_hooks_stop_on_first_error() {
409 let calls = Arc::new(Mutex::new(Vec::new()));
410 let chain = RuntimeHookChain::new(vec![
411 hook(RecordingHook::new("first", calls.clone()).with_action_suffix(1)),
412 hook(RecordingHook::new("second", calls.clone()).with_transform_error("bad action")),
413 hook(RecordingHook::new("third", calls.clone()).with_action_suffix(3)),
414 ]);
415
416 let error = chain
417 .transform_action(action_event(vec![0]))
418 .await
419 .unwrap_err();
420
421 assert_eq!(error.to_string(), "bad action");
422 assert_eq!(
423 recorded(&calls),
424 vec!["first:action:[0]", "second:action:[0, 1]"]
425 );
426 }
427
428 #[tokio::test]
429 async fn empty_chain_is_a_noop() {
430 let chain = RuntimeHookChain::empty();
431
432 chain.log(log_event()).await.unwrap();
433 let action = chain
434 .transform_action(action_event(vec![7]))
435 .await
436 .unwrap()
437 .unwrap();
438
439 assert!(chain.is_empty());
440 assert_eq!(chain.len(), 0);
441 assert_eq!(action.data, vec![7]);
442 }
443}