use async_trait::async_trait;
use prost::bytes::Bytes;
use super::{
ActionReceivedEvent, EnvConnectedEvent, EpisodeCompletedEvent, EpisodeStartedEvent, LogEvent,
ModelConnectedEvent, ObservationEmittedEvent, SessionEndedEvent, SessionFailedEvent,
SessionStartedEvent, StepCompletedEvent, TelemetrySnapshotEvent,
};
#[derive(Debug, thiserror::Error)]
#[non_exhaustive]
pub enum HookError {
#[error("{0}")]
Message(String),
}
#[derive(Debug, Default)]
pub struct NoopRuntimeHooks;
#[async_trait]
impl RuntimeHooks for NoopRuntimeHooks {}
#[async_trait]
pub trait RuntimeHooks: Send + Sync {
async fn env_connected(&self, _event: EnvConnectedEvent) -> Result<(), HookError> {
Ok(())
}
async fn model_connected(&self, _event: ModelConnectedEvent) -> Result<(), HookError> {
Ok(())
}
async fn session_started(&self, _event: SessionStartedEvent) -> Result<(), HookError> {
Ok(())
}
async fn episode_started(&self, _event: EpisodeStartedEvent) -> Result<(), HookError> {
Ok(())
}
async fn episode_completed(&self, _event: EpisodeCompletedEvent) -> Result<(), HookError> {
Ok(())
}
async fn action_received(&self, _event: ActionReceivedEvent) -> Result<(), HookError> {
Ok(())
}
async fn transform_action(
&self,
event: ActionReceivedEvent,
) -> Result<Option<Vec<Bytes>>, HookError> {
Ok(event.action)
}
async fn step_completed(&self, _event: StepCompletedEvent) -> Result<(), HookError> {
Ok(())
}
async fn observation_emitted(&self, _event: ObservationEmittedEvent) -> Result<(), HookError> {
Ok(())
}
async fn transform_observation(
&self,
event: ObservationEmittedEvent,
) -> Result<Option<Vec<Bytes>>, HookError> {
Ok(event.observation)
}
async fn session_ended(&self, _event: SessionEndedEvent) -> Result<(), HookError> {
Ok(())
}
async fn on_telemetry(&self, _event: TelemetrySnapshotEvent) -> Result<(), HookError> {
Ok(())
}
async fn session_failed(&self, _event: SessionFailedEvent) -> Result<(), HookError> {
Ok(())
}
async fn log(&self, _event: LogEvent) -> Result<(), HookError> {
Ok(())
}
}