use std::sync::Arc;
use async_trait::async_trait;
use super::types::ModelObservation;
use crate::{Result, spaces};
#[async_trait]
pub trait ModelRouteSetup: Send + Sync {
async fn resolve_adapter(
&self,
env_id: &str,
env_contract: &spaces::EnvContract,
options: ResolveOptions,
) -> Result<RouteNeeds>;
async fn reset_adapter(&self, _env_id: &str, _episode_ids: &[String]) -> Result<()> {
Ok(())
}
async fn release_adapter(&self, _env_id: &str) -> Result<()> {
Ok(())
}
}
#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
pub struct ResolveOptions {
pub execution_horizon: u32,
pub delivers_history: bool,
}
#[derive(Clone, Debug, Default, PartialEq, Eq)]
pub struct RouteNeeds {
pub native_chunk: Option<u32>,
pub history: Option<HistoryNeeds>,
}
#[derive(Clone, Debug, Default, PartialEq, Eq)]
pub struct HistoryNeeds {
pub keys: Vec<String>,
pub prunable: bool,
}
#[derive(Debug)]
pub struct PredictFrames {
pub actions: Vec<spaces::SpaceValue>,
pub replay: Vec<Vec<spaces::SpaceValue>>,
}
#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
pub struct HeldState {
pub episodes: u64,
pub bytes: u64,
}
#[async_trait]
pub trait ModelHandler: Send {
async fn predict(&mut self, observation: ModelObservation) -> Result<Vec<spaces::SpaceValue>>;
async fn predict_chunked(&mut self, observation: ModelObservation) -> Result<PredictFrames> {
Ok(PredictFrames {
actions: self.predict(observation).await?,
replay: Vec::new(),
})
}
async fn predict_grouped(
&mut self,
observations: Vec<ModelObservation>,
) -> Vec<Result<PredictFrames>> {
let mut results = Vec::with_capacity(observations.len());
for observation in observations {
results.push(self.predict_chunked(observation).await);
}
results
}
fn held_state(&self) -> Option<HeldState> {
None
}
fn take_adapter_ns(&mut self) -> u64 {
0
}
fn route_setup(&self) -> Option<Arc<dyn ModelRouteSetup>> {
None
}
async fn reset_adapter(&mut self, _env_id: &str, _episode_ids: Vec<String>) -> Result<()> {
Ok(())
}
async fn on_close(&mut self) -> Result<()> {
Ok(())
}
}