Skip to main content

rlmesh_runtime/hooks/
chain.rs

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/// Ordered runtime hook composition.
12///
13/// Lifecycle, progress, telemetry, and log events are sent to every hook.
14/// Transform hooks are applied in order, with each hook receiving the payload
15/// returned by the previous hook.
16#[derive(Default)]
17pub struct RuntimeHookChain {
18    hooks: Vec<std::sync::Arc<dyn RuntimeHooks>>,
19}
20
21impl RuntimeHookChain {
22    /// Creates a chain from hooks in invocation order.
23    pub fn new(hooks: Vec<std::sync::Arc<dyn RuntimeHooks>>) -> Self {
24        Self { hooks }
25    }
26
27    /// Creates an empty chain with the same behavior as `NoopRuntimeHooks`.
28    pub fn empty() -> Self {
29        Self::default()
30    }
31
32    /// Returns the number of hooks in the chain.
33    pub fn len(&self) -> usize {
34        self.hooks.len()
35    }
36
37    /// Returns true when the chain contains no hooks.
38    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}