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, Eq)]
pub struct RunLocalOptions {
pub env_address: ConnectAddress,
pub max_episodes: Option<u64>,
pub base_seed: Option<i64>,
}
impl RunLocalOptions {
pub fn new(env_address: ConnectAddress) -> Self {
Self {
env_address,
max_episodes: None,
base_seed: None,
}
}
pub fn parse(env_address: &str) -> Result<Self> {
Ok(Self::new(ConnectAddress::parse(env_address)?))
}
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
}
}
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(
mut self,
options: impl Into<RunLocalOptions>,
) -> Result<RuntimeReport> {
let options = options.into();
let result = local::run_local(
&mut self.handler,
options.env_address,
options.max_episodes,
options.base_seed,
)
.await;
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();
server::bind_model_with_options(
self.handler,
options.address,
&options.token,
options.serve,
)
.await
}
}