use std::sync::Arc;
use rlmesh_runtime::RuntimeReport;
use super::handler::ModelHandler;
use super::server::BoundModelServer;
use super::{local, server};
use crate::{BindAddress, ConnectAddress, Error, Result, ServeOptions};
pub struct ModelWorker<H> {
handler: H,
}
impl<H> ModelWorker<H> {
pub fn new(handler: H) -> Self {
Self { handler }
}
}
#[derive(Debug, Clone, PartialEq)]
pub struct RunLocalOptions {
pub env_address: ConnectAddress,
pub token: String,
pub max_episodes: Option<u64>,
pub base_seed: Option<i64>,
pub episode_seeds: Vec<i64>,
pub trial_index_base: Option<u64>,
pub max_episode_steps: Option<i64>,
pub max_episode_seconds: Option<f64>,
pub close_env: bool,
pub close_model: bool,
pub execution_horizon: u32,
pub prefetch_lead: u32,
pub workflow_edition: Option<String>,
}
impl RunLocalOptions {
pub fn new(env_address: ConnectAddress) -> Self {
Self {
env_address,
token: String::new(),
max_episodes: None,
base_seed: None,
episode_seeds: Vec::new(),
trial_index_base: None,
max_episode_steps: None,
max_episode_seconds: None,
close_env: false,
close_model: true,
execution_horizon: 1,
prefetch_lead: 0,
workflow_edition: None,
}
}
pub fn parse(env_address: &str) -> Result<Self> {
Ok(Self::new(ConnectAddress::parse(env_address)?))
}
pub fn token(mut self, token: impl Into<String>) -> Self {
self.token = token.into();
self
}
pub fn for_episodes(mut self, max_episodes: u64) -> Self {
self.max_episodes = Some(max_episodes);
self
}
pub fn base_seed(mut self, base_seed: i64) -> Self {
self.base_seed = Some(base_seed);
self
}
pub fn execution_horizon(mut self, execution_horizon: u32) -> Self {
self.execution_horizon = execution_horizon.max(1);
self
}
pub fn prefetch_lead(mut self, prefetch_lead: u32) -> Self {
self.prefetch_lead = prefetch_lead;
self
}
pub fn episode_seeds(mut self, episode_seeds: Vec<i64>) -> Self {
self.episode_seeds = episode_seeds;
self
}
pub fn trial_index_base(mut self, trial_index_base: u64) -> Self {
self.trial_index_base = Some(trial_index_base);
self
}
pub fn max_episode_steps(mut self, max_episode_steps: i64) -> Self {
self.max_episode_steps = Some(max_episode_steps);
self
}
pub fn max_episode_seconds(mut self, max_episode_seconds: f64) -> Self {
self.max_episode_seconds = Some(max_episode_seconds);
self
}
pub fn close_env(mut self, close_env: bool) -> Self {
self.close_env = close_env;
self
}
pub fn close_model(mut self, close_model: bool) -> Self {
self.close_model = close_model;
self
}
pub fn workflow_edition(mut self, workflow_edition: impl Into<String>) -> Self {
self.workflow_edition = Some(workflow_edition.into());
self
}
}
impl From<ConnectAddress> for RunLocalOptions {
fn from(env_address: ConnectAddress) -> Self {
Self::new(env_address)
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct ServeModelOptions {
pub address: BindAddress,
pub token: String,
pub serve: ServeOptions,
}
impl ServeModelOptions {
pub fn new(address: BindAddress) -> Self {
Self {
address,
token: String::new(),
serve: ServeOptions::default(),
}
}
pub fn parse(address: &str) -> Result<Self> {
Ok(Self::new(BindAddress::parse(address)?))
}
pub fn token(mut self, token: impl Into<String>) -> Self {
self.token = token.into();
self
}
pub fn serve_options(mut self, serve: ServeOptions) -> Self {
self.serve = serve;
self
}
}
impl From<BindAddress> for ServeModelOptions {
fn from(address: BindAddress) -> Self {
Self::new(address)
}
}
impl<H: ModelHandler + 'static> ModelWorker<H> {
pub fn run_local(self, options: impl Into<RunLocalOptions>) -> Result<RuntimeReport> {
let runtime = tokio::runtime::Runtime::new()
.map_err(|err| Error::Internal(format!("failed to create tokio runtime: {err}")))?;
runtime.block_on(self.run_local_async(options))
}
pub async fn run_local_async(
self,
options: impl Into<RunLocalOptions>,
) -> Result<RuntimeReport> {
self.run_local_cancellable_async(options, tokio_util::sync::CancellationToken::new())
.await
}
pub async fn run_local_cancellable_async(
self,
options: impl Into<RunLocalOptions>,
cancellation: tokio_util::sync::CancellationToken,
) -> Result<RuntimeReport> {
self.run_local_hooked_async(
options,
cancellation,
Arc::new(rlmesh_runtime::NoopRuntimeHooks),
)
.await
}
pub async fn run_local_hooked_async(
mut self,
options: impl Into<RunLocalOptions>,
cancellation: tokio_util::sync::CancellationToken,
hooks: Arc<dyn rlmesh_runtime::RuntimeHooks>,
) -> Result<RuntimeReport> {
let options = options.into();
let close_model = options.close_model;
let result = local::run_local(&mut self.handler, options, cancellation, hooks).await;
if !close_model {
return result;
}
let close_result = self.handler.on_close().await;
crate::error::join_results(result, close_result, "local model run failed")
}
pub fn serve(self, options: impl Into<ServeModelOptions>) -> Result<()> {
let runtime = tokio::runtime::Runtime::new()
.map_err(|err| Error::Internal(format!("failed to create tokio runtime: {err}")))?;
runtime.block_on(self.serve_async(options))
}
pub async fn serve_async(self, options: impl Into<ServeModelOptions>) -> Result<()> {
self.bind_async(options).await?.serve().await
}
pub async fn bind_async(
self,
options: impl Into<ServeModelOptions>,
) -> Result<BoundModelServer> {
let options = options.into();
let effective_token = options
.serve
.token
.clone()
.filter(|token| !token.is_empty())
.unwrap_or_else(|| options.token.clone());
server::bind_model_with_options(
self.handler,
options.address,
&effective_token,
options.serve,
)
.await
}
}