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,
execution_horizon: u32,
) -> Result<()>;
async fn release_adapter(&self, _env_id: &str) -> Result<()> {
Ok(())
}
}
pub struct PredictFrames {
pub actions: Vec<spaces::SpaceValue>,
pub replay: Vec<Vec<spaces::SpaceValue>>,
}
#[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<Vec<spaces::SpaceValue>>> {
let mut results = Vec::with_capacity(observations.len());
for observation in observations {
results.push(self.predict(observation).await);
}
results
}
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(())
}
}