Skip to main content

rlmesh_runtime/
spec.rs

1use std::time::Duration;
2
3use rlmesh_proto::env::v1::EnvContract;
4use rlmesh_proto::spaces::v1::SpaceSpec;
5use serde::{Deserialize, Serialize};
6
7#[derive(Debug, Clone, PartialEq)]
8pub struct RuntimeSessionSpec {
9    pub session_id: String,
10    pub route_id: String,
11    pub env_component_id: String,
12    pub model_component_id: String,
13    pub env_id: String,
14    pub env_contract: EnvContract,
15    pub num_envs: usize,
16    pub max_episodes: Option<u64>,
17    pub close_env_on_end: bool,
18    pub limits: RuntimeLimits,
19}
20
21impl RuntimeSessionSpec {
22    pub fn validate(&self) -> Result<(), String> {
23        if self.session_id.trim().is_empty() {
24            return Err("runtime session_id must not be empty".to_string());
25        }
26        if self.route_id.trim().is_empty() {
27            return Err("runtime route_id must not be empty".to_string());
28        }
29        if self.env_component_id.trim().is_empty() {
30            return Err("runtime env_component_id must not be empty".to_string());
31        }
32        if self.model_component_id.trim().is_empty() {
33            return Err("runtime model_component_id must not be empty".to_string());
34        }
35        if self.num_envs == 0 {
36            return Err("runtime num_envs must be greater than zero".to_string());
37        }
38        if self.env_contract.observation_space.is_none() {
39            return Err("runtime env_contract is missing observation_space".to_string());
40        }
41        if self.env_contract.action_space.is_none() {
42            return Err("runtime env_contract is missing action_space".to_string());
43        }
44        if self.max_episodes == Some(0) {
45            return Err("runtime max_episodes must be greater than zero when set".to_string());
46        }
47        Ok(())
48    }
49
50    pub fn route_context(&self) -> crate::hooks::RuntimeRouteContext {
51        crate::hooks::RuntimeRouteContext {
52            route_id: self.route_id.clone(),
53            env_component_id: self.env_component_id.clone(),
54            model_component_id: self.model_component_id.clone(),
55        }
56    }
57
58    pub fn observation_space(&self) -> &SpaceSpec {
59        self.env_contract
60            .observation_space
61            .as_ref()
62            .expect("RuntimeSessionSpec validated observation_space")
63    }
64
65    pub fn action_space(&self) -> &SpaceSpec {
66        self.env_contract
67            .action_space
68            .as_ref()
69            .expect("RuntimeSessionSpec validated action_space")
70    }
71}
72
73#[derive(Debug, Clone, PartialEq, Eq)]
74pub struct RuntimeReport {
75    pub session_id: String,
76    pub route_id: String,
77    pub total_steps: i64,
78    pub total_episodes: i64,
79}
80
81#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
82#[serde(rename_all = "camelCase", deny_unknown_fields)]
83pub struct RuntimeLimits {
84    #[serde(
85        default = "default_connect_timeout",
86        rename = "envConnectTimeoutMs",
87        serialize_with = "duration_millis::serialize",
88        deserialize_with = "duration_millis::deserialize"
89    )]
90    pub env_connect_timeout: Duration,
91    #[serde(
92        default = "default_model_connect_timeout",
93        rename = "modelConnectTimeoutMs",
94        serialize_with = "duration_millis::serialize",
95        deserialize_with = "duration_millis::deserialize"
96    )]
97    pub model_connect_timeout: Duration,
98    #[serde(
99        default = "default_configure_route_timeout",
100        rename = "configureRouteTimeoutMs",
101        serialize_with = "duration_millis::serialize",
102        deserialize_with = "duration_millis::deserialize"
103    )]
104    pub configure_route_timeout: Duration,
105    #[serde(
106        default = "default_env_reset_timeout",
107        rename = "envResetTimeoutMs",
108        serialize_with = "duration_millis::serialize",
109        deserialize_with = "duration_millis::deserialize"
110    )]
111    pub env_reset_timeout: Duration,
112    #[serde(
113        default = "default_model_predict_timeout",
114        rename = "modelPredictTimeoutMs",
115        serialize_with = "duration_millis::serialize",
116        deserialize_with = "duration_millis::deserialize"
117    )]
118    pub model_predict_timeout: Duration,
119    #[serde(
120        default = "default_env_step_timeout",
121        rename = "envStepTimeoutMs",
122        serialize_with = "duration_millis::serialize",
123        deserialize_with = "duration_millis::deserialize"
124    )]
125    pub env_step_timeout: Duration,
126    #[serde(
127        default = "default_service_close_timeout",
128        rename = "serviceCloseTimeoutMs",
129        serialize_with = "duration_millis::serialize",
130        deserialize_with = "duration_millis::deserialize"
131    )]
132    pub service_close_timeout: Duration,
133    #[serde(
134        default = "default_telemetry_window",
135        rename = "telemetryWindowMs",
136        serialize_with = "duration_millis::serialize",
137        deserialize_with = "duration_millis::deserialize"
138    )]
139    pub telemetry_window: Duration,
140}
141
142impl Default for RuntimeLimits {
143    fn default() -> Self {
144        Self {
145            env_connect_timeout: default_connect_timeout(),
146            model_connect_timeout: default_model_connect_timeout(),
147            configure_route_timeout: default_configure_route_timeout(),
148            env_reset_timeout: default_env_reset_timeout(),
149            model_predict_timeout: default_model_predict_timeout(),
150            env_step_timeout: default_env_step_timeout(),
151            service_close_timeout: default_service_close_timeout(),
152            telemetry_window: default_telemetry_window(),
153        }
154    }
155}
156
157impl RuntimeLimits {
158    pub fn env_step_timeout_ms(&self) -> i64 {
159        duration_ms_i64(self.env_step_timeout)
160    }
161
162    pub fn env_reset_timeout_ms(&self) -> i64 {
163        duration_ms_i64(self.env_reset_timeout)
164    }
165}
166
167fn default_connect_timeout() -> Duration {
168    Duration::from_secs(60)
169}
170
171fn default_model_connect_timeout() -> Duration {
172    Duration::from_secs(600)
173}
174
175fn default_configure_route_timeout() -> Duration {
176    Duration::from_secs(600)
177}
178
179fn default_env_reset_timeout() -> Duration {
180    Duration::from_secs(300)
181}
182
183fn default_model_predict_timeout() -> Duration {
184    Duration::from_secs(300)
185}
186
187fn default_env_step_timeout() -> Duration {
188    Duration::from_secs(300)
189}
190
191fn default_service_close_timeout() -> Duration {
192    Duration::from_secs(5)
193}
194
195fn default_telemetry_window() -> Duration {
196    Duration::from_secs(1)
197}
198
199fn duration_ms_i64(duration: Duration) -> i64 {
200    duration.as_millis().try_into().unwrap_or(i64::MAX)
201}
202
203mod duration_millis {
204    use std::time::Duration;
205
206    use serde::{Deserialize, Deserializer, Serializer};
207
208    pub(super) fn serialize<S>(duration: &Duration, serializer: S) -> Result<S::Ok, S::Error>
209    where
210        S: Serializer,
211    {
212        let millis = duration.as_millis().try_into().unwrap_or(u64::MAX);
213        serializer.serialize_u64(millis)
214    }
215
216    pub(super) fn deserialize<'de, D>(deserializer: D) -> Result<Duration, D::Error>
217    where
218        D: Deserializer<'de>,
219    {
220        Ok(Duration::from_millis(u64::deserialize(deserializer)?))
221    }
222}
223
224#[cfg(test)]
225mod tests {
226    use std::time::Duration;
227
228    use serde_json::json;
229
230    use super::RuntimeLimits;
231
232    #[test]
233    fn runtime_limits_json_uses_explicit_millisecond_fields() {
234        let value = serde_json::to_value(RuntimeLimits::default()).unwrap();
235
236        assert_eq!(value["envConnectTimeoutMs"], json!(60_000));
237        assert_eq!(value["modelConnectTimeoutMs"], json!(600_000));
238        assert_eq!(value["configureRouteTimeoutMs"], json!(600_000));
239        assert_eq!(value["envResetTimeoutMs"], json!(300_000));
240        assert_eq!(value["modelPredictTimeoutMs"], json!(300_000));
241        assert_eq!(value["envStepTimeoutMs"], json!(300_000));
242        assert_eq!(value["serviceCloseTimeoutMs"], json!(5_000));
243        assert_eq!(value["telemetryWindowMs"], json!(1_000));
244        assert!(value.get("envConnectTimeout").is_none());
245        assert!(value.get("telemetryWindow").is_none());
246
247        let parsed: RuntimeLimits = serde_json::from_value(json!({
248            "envConnectTimeoutMs": 1,
249            "modelConnectTimeoutMs": 2,
250            "configureRouteTimeoutMs": 3,
251            "envResetTimeoutMs": 4,
252            "modelPredictTimeoutMs": 5,
253            "envStepTimeoutMs": 6,
254            "serviceCloseTimeoutMs": 7,
255            "telemetryWindowMs": 8
256        }))
257        .unwrap();
258
259        assert_eq!(parsed.env_connect_timeout, Duration::from_millis(1));
260        assert_eq!(parsed.model_connect_timeout, Duration::from_millis(2));
261        assert_eq!(parsed.configure_route_timeout, Duration::from_millis(3));
262        assert_eq!(parsed.env_reset_timeout, Duration::from_millis(4));
263        assert_eq!(parsed.model_predict_timeout, Duration::from_millis(5));
264        assert_eq!(parsed.env_step_timeout, Duration::from_millis(6));
265        assert_eq!(parsed.service_close_timeout, Duration::from_millis(7));
266        assert_eq!(parsed.telemetry_window, Duration::from_millis(8));
267    }
268
269    #[test]
270    fn runtime_limits_reject_legacy_unsuffixed_fields() {
271        let error = serde_json::from_value::<RuntimeLimits>(json!({
272            "envConnectTimeout": 1
273        }))
274        .unwrap_err();
275
276        assert!(error.to_string().contains("envConnectTimeout"));
277    }
278}