use std::sync::Arc;
use prost::bytes::Bytes;
use rlmesh_proto::spaces::v1::{SpaceSpec, space_spec::Spec};
use rlmesh_spaces::{Advisory, DType};
use crate::spec::PeerCeiling;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum Leg {
EnvToModel,
ModelToEnv,
}
impl Leg {
pub fn target(self) -> &'static str {
match self {
Leg::EnvToModel => "model",
Leg::ModelToEnv => "env",
}
}
}
#[derive(Debug, Clone, PartialEq)]
pub struct PayloadFacts {
pub space: Arc<SpaceSpec>,
pub byte_len: usize,
pub leaves: Vec<Bytes>,
}
impl PayloadFacts {
pub fn new(space: Arc<SpaceSpec>, leaves: Vec<Bytes>) -> Self {
Self {
space,
byte_len: leaves.iter().map(Bytes::len).sum(),
leaves,
}
}
pub fn exceeds(&self, ceiling: &PeerCeiling) -> Option<String> {
if let Some(dtype) = dtype_outside(&self.space, ceiling) {
return Some(format!("dtype {} is outside its ceiling", dtype.name()));
}
if self.byte_len > ceiling.max_message_size {
return Some(format!(
"{} bytes exceeds its {}-byte message cap",
self.byte_len, ceiling.max_message_size
));
}
None
}
}
fn dtype_outside(space: &SpaceSpec, ceiling: &PeerCeiling) -> Option<DType> {
match &space.spec {
Some(Spec::Dict(dict)) => dict
.spaces
.iter()
.find_map(|space| dtype_outside(space, ceiling)),
Some(Spec::Tuple(tuple)) => tuple
.spaces
.iter()
.find_map(|space| dtype_outside(space, ceiling)),
_ => DType::try_from(space.dtype)
.ok()
.filter(|dtype| !ceiling.dtypes.contains(dtype)),
}
}
#[derive(Debug, Clone, PartialEq)]
pub enum RelayDecision {
Forward,
Convert {
leaves: Vec<Bytes>,
space: Option<Arc<SpaceSpec>>,
advisory: Advisory,
},
Refuse(String),
}
pub trait RelayPolicy: Send + Sync {
fn reconcile(&self, leg: Leg, ceiling: &PeerCeiling, payload: &PayloadFacts) -> RelayDecision;
}
#[derive(Debug, Default)]
pub struct RefusingRelayPolicy;
impl RelayPolicy for RefusingRelayPolicy {
fn reconcile(&self, _leg: Leg, ceiling: &PeerCeiling, payload: &PayloadFacts) -> RelayDecision {
payload
.exceeds(ceiling)
.map_or(RelayDecision::Forward, RelayDecision::Refuse)
}
}