use async_trait::async_trait;
use rlmesh_proto::common::v1::MessageBytes;
use super::{
ActionReceivedEvent, EnvConnectedEvent, EpisodeCompletedEvent, EpisodeStartedEvent, HookError,
LogEvent, ModelConnectedEvent, ObservationEmittedEvent, RuntimeHooks, SessionEndedEvent,
SessionFailedEvent, SessionStartedEvent, StepCompletedEvent, TelemetrySummaryEvent,
TelemetryWindowEvent,
};
#[derive(Default)]
pub struct RuntimeHookChain {
hooks: Vec<std::sync::Arc<dyn RuntimeHooks>>,
}
impl RuntimeHookChain {
pub fn new(hooks: Vec<std::sync::Arc<dyn RuntimeHooks>>) -> Self {
Self { hooks }
}
pub fn empty() -> Self {
Self::default()
}
pub fn len(&self) -> usize {
self.hooks.len()
}
pub fn is_empty(&self) -> bool {
self.hooks.is_empty()
}
}
#[async_trait]
impl RuntimeHooks for RuntimeHookChain {
async fn env_connected(&self, event: EnvConnectedEvent) -> Result<(), HookError> {
let mut first_error = None;
for hook in &self.hooks {
if let Err(error) = hook.env_connected(event.clone()).await {
first_error.get_or_insert(error);
}
}
first_error.map_or(Ok(()), Err)
}
async fn model_connected(&self, event: ModelConnectedEvent) -> Result<(), HookError> {
let mut first_error = None;
for hook in &self.hooks {
if let Err(error) = hook.model_connected(event.clone()).await {
first_error.get_or_insert(error);
}
}
first_error.map_or(Ok(()), Err)
}
async fn session_started(&self, event: SessionStartedEvent) -> Result<(), HookError> {
let mut first_error = None;
for hook in &self.hooks {
if let Err(error) = hook.session_started(event.clone()).await {
first_error.get_or_insert(error);
}
}
first_error.map_or(Ok(()), Err)
}
async fn episode_started(&self, event: EpisodeStartedEvent) -> Result<(), HookError> {
let mut first_error = None;
for hook in &self.hooks {
if let Err(error) = hook.episode_started(event.clone()).await {
first_error.get_or_insert(error);
}
}
first_error.map_or(Ok(()), Err)
}
async fn episode_completed(&self, event: EpisodeCompletedEvent) -> Result<(), HookError> {
let mut first_error = None;
for hook in &self.hooks {
if let Err(error) = hook.episode_completed(event.clone()).await {
first_error.get_or_insert(error);
}
}
first_error.map_or(Ok(()), Err)
}
async fn action_received(&self, event: ActionReceivedEvent) -> Result<(), HookError> {
let mut first_error = None;
for hook in &self.hooks {
if let Err(error) = hook.action_received(event.clone()).await {
first_error.get_or_insert(error);
}
}
first_error.map_or(Ok(()), Err)
}
async fn transform_action(
&self,
event: ActionReceivedEvent,
) -> Result<Option<MessageBytes>, HookError> {
let ActionReceivedEvent {
session_id,
route,
episode_id,
episode_record_id,
episode_ids,
episode_record_ids,
step,
env_index,
action_space,
mut action,
} = event;
for hook in &self.hooks {
action = hook
.transform_action(ActionReceivedEvent {
session_id: session_id.clone(),
route: route.clone(),
episode_id: episode_id.clone(),
episode_record_id: episode_record_id.clone(),
episode_ids: episode_ids.clone(),
episode_record_ids: episode_record_ids.clone(),
step,
env_index,
action_space: action_space.clone(),
action,
})
.await?;
}
Ok(action)
}
async fn step_completed(&self, event: StepCompletedEvent) -> Result<(), HookError> {
let mut first_error = None;
for hook in &self.hooks {
if let Err(error) = hook.step_completed(event.clone()).await {
first_error.get_or_insert(error);
}
}
first_error.map_or(Ok(()), Err)
}
async fn observation_emitted(&self, event: ObservationEmittedEvent) -> Result<(), HookError> {
let mut first_error = None;
for hook in &self.hooks {
if let Err(error) = hook.observation_emitted(event.clone()).await {
first_error.get_or_insert(error);
}
}
first_error.map_or(Ok(()), Err)
}
async fn transform_observation(
&self,
event: ObservationEmittedEvent,
) -> Result<Option<MessageBytes>, HookError> {
let ObservationEmittedEvent {
session_id,
route,
episode_id,
episode_record_id,
episode_ids,
episode_record_ids,
step,
env_index,
is_reset,
num_envs,
observation_space,
mut observation,
} = event;
for hook in &self.hooks {
observation = hook
.transform_observation(ObservationEmittedEvent {
session_id: session_id.clone(),
route: route.clone(),
episode_id: episode_id.clone(),
episode_record_id: episode_record_id.clone(),
episode_ids: episode_ids.clone(),
episode_record_ids: episode_record_ids.clone(),
step,
env_index,
is_reset,
num_envs,
observation_space: observation_space.clone(),
observation,
})
.await?;
}
Ok(observation)
}
async fn telemetry_window(&self, event: TelemetryWindowEvent) -> Result<(), HookError> {
let mut first_error = None;
for hook in &self.hooks {
if let Err(error) = hook.telemetry_window(event.clone()).await {
first_error.get_or_insert(error);
}
}
first_error.map_or(Ok(()), Err)
}
async fn telemetry_summary(&self, event: TelemetrySummaryEvent) -> Result<(), HookError> {
let mut first_error = None;
for hook in &self.hooks {
if let Err(error) = hook.telemetry_summary(event.clone()).await {
first_error.get_or_insert(error);
}
}
first_error.map_or(Ok(()), Err)
}
async fn session_ended(&self, event: SessionEndedEvent) -> Result<(), HookError> {
let mut first_error = None;
for hook in &self.hooks {
if let Err(error) = hook.session_ended(event.clone()).await {
first_error.get_or_insert(error);
}
}
first_error.map_or(Ok(()), Err)
}
async fn session_failed(&self, event: SessionFailedEvent) -> Result<(), HookError> {
let mut first_error = None;
for hook in &self.hooks {
if let Err(error) = hook.session_failed(event.clone()).await {
first_error.get_or_insert(error);
}
}
first_error.map_or(Ok(()), Err)
}
async fn log(&self, event: LogEvent) -> Result<(), HookError> {
let mut first_error = None;
for hook in &self.hooks {
if let Err(error) = hook.log(event.clone()).await {
first_error.get_or_insert(error);
}
}
first_error.map_or(Ok(()), Err)
}
}
#[cfg(test)]
mod tests {
use std::sync::{Arc, Mutex};
use async_trait::async_trait;
use rlmesh_proto::common::v1::MessageBytes;
use rlmesh_proto::spaces::v1::SpaceSpec;
use super::*;
use crate::hooks::{LogLevel, RuntimeRouteContext};
struct RecordingHook {
name: &'static str,
calls: Arc<Mutex<Vec<String>>>,
log_error: Option<&'static str>,
action_suffix: Option<u8>,
transform_error: Option<&'static str>,
}
impl RecordingHook {
fn new(name: &'static str, calls: Arc<Mutex<Vec<String>>>) -> Self {
Self {
name,
calls,
log_error: None,
action_suffix: None,
transform_error: None,
}
}
fn with_log_error(mut self, error: &'static str) -> Self {
self.log_error = Some(error);
self
}
fn with_action_suffix(mut self, suffix: u8) -> Self {
self.action_suffix = Some(suffix);
self
}
fn with_transform_error(mut self, error: &'static str) -> Self {
self.transform_error = Some(error);
self
}
fn record(&self, call: impl Into<String>) {
self.calls
.lock()
.expect("calls mutex poisoned")
.push(call.into());
}
}
#[async_trait]
impl RuntimeHooks for RecordingHook {
async fn log(&self, event: LogEvent) -> Result<(), HookError> {
self.record(format!("{}:log:{}", self.name, event.message));
if let Some(error) = self.log_error {
return Err(HookError::Message(error.to_string()));
}
Ok(())
}
async fn transform_action(
&self,
event: ActionReceivedEvent,
) -> Result<Option<MessageBytes>, HookError> {
let data = event
.action
.as_ref()
.map(|action| action.data.clone())
.unwrap_or_default();
self.record(format!("{}:action:{data:?}", self.name));
if let Some(error) = self.transform_error {
return Err(HookError::Message(error.to_string()));
}
Ok(event.action.map(|mut action| {
if let Some(suffix) = self.action_suffix {
action.data.push(suffix);
}
action
}))
}
}
fn hook(hook: RecordingHook) -> Arc<dyn RuntimeHooks> {
Arc::new(hook)
}
fn recorded(calls: &Arc<Mutex<Vec<String>>>) -> Vec<String> {
calls.lock().expect("calls mutex poisoned").clone()
}
fn log_event() -> LogEvent {
LogEvent {
session_id: "session".to_string(),
route: RuntimeRouteContext::default(),
level: LogLevel::Info,
message: "hello".to_string(),
source: None,
}
}
fn action_event(data: Vec<u8>) -> ActionReceivedEvent {
ActionReceivedEvent {
session_id: "session".to_string(),
route: RuntimeRouteContext::default(),
episode_id: "episode".to_string(),
episode_record_id: "episode-artifact".to_string(),
episode_ids: vec!["episode".to_string()],
episode_record_ids: vec!["episode-artifact".to_string()],
step: 1,
env_index: 0,
action_space: SpaceSpec::default(),
action: Some(MessageBytes { data }),
}
}
#[tokio::test]
async fn event_hooks_call_every_hook_and_return_first_error() {
let calls = Arc::new(Mutex::new(Vec::new()));
let chain = RuntimeHookChain::new(vec![
hook(RecordingHook::new("first", calls.clone()).with_log_error("first failed")),
hook(RecordingHook::new("second", calls.clone()).with_log_error("second failed")),
hook(RecordingHook::new("third", calls.clone())),
]);
let error = chain.log(log_event()).await.unwrap_err();
assert_eq!(error.to_string(), "first failed");
assert_eq!(
recorded(&calls),
vec!["first:log:hello", "second:log:hello", "third:log:hello"]
);
}
#[tokio::test]
async fn transform_hooks_run_in_order() {
let calls = Arc::new(Mutex::new(Vec::new()));
let chain = RuntimeHookChain::new(vec![
hook(RecordingHook::new("first", calls.clone()).with_action_suffix(1)),
hook(RecordingHook::new("second", calls.clone()).with_action_suffix(2)),
]);
let action = chain
.transform_action(action_event(vec![0]))
.await
.unwrap()
.unwrap();
assert_eq!(action.data, vec![0, 1, 2]);
assert_eq!(
recorded(&calls),
vec!["first:action:[0]", "second:action:[0, 1]"]
);
}
#[tokio::test]
async fn transform_hooks_stop_on_first_error() {
let calls = Arc::new(Mutex::new(Vec::new()));
let chain = RuntimeHookChain::new(vec![
hook(RecordingHook::new("first", calls.clone()).with_action_suffix(1)),
hook(RecordingHook::new("second", calls.clone()).with_transform_error("bad action")),
hook(RecordingHook::new("third", calls.clone()).with_action_suffix(3)),
]);
let error = chain
.transform_action(action_event(vec![0]))
.await
.unwrap_err();
assert_eq!(error.to_string(), "bad action");
assert_eq!(
recorded(&calls),
vec!["first:action:[0]", "second:action:[0, 1]"]
);
}
#[tokio::test]
async fn empty_chain_is_a_noop() {
let chain = RuntimeHookChain::empty();
chain.log(log_event()).await.unwrap();
let action = chain
.transform_action(action_event(vec![7]))
.await
.unwrap()
.unwrap();
assert!(chain.is_empty());
assert_eq!(chain.len(), 0);
assert_eq!(action.data, vec![7]);
}
}