1use std::cmp::max;
2use std::cmp::min;
3use std::collections::HashMap;
4use std::thread;
5use std::time::Duration;
6
7use itertools::izip;
8use itertools::Itertools;
9use pyo3::exceptions::asyncio::InvalidStateError;
10use pyo3::intern;
11use pyo3::prelude::*;
12use pyo3::sync::GILOnceCell;
13use pyo3::types::PyDict;
14use pyo3::IntoPyObjectExt;
15use raw_sync::events::Event;
16use raw_sync::events::EventInit;
17use raw_sync::events::EventState;
18use shared_memory::Shmem;
19use shared_memory::ShmemConf;
20
21use crate::common::misc::clone_list;
22use crate::communication::retrieve_python;
23use crate::communication::{
24 append_header, get_flink, recvfrom_byte, retrieve_bool, retrieve_usize, sendto_byte, Header,
25};
26use crate::env_action::append_env_action_new;
27use crate::env_action::EnvAction;
28use crate::serdes::pyany_serde::PythonSerde;
29
30fn sync_with_env_process<'py>(
31 py: Python<'py>,
32 socket: &PyObject,
33 address: &PyObject,
34) -> PyResult<()> {
35 recvfrom_byte(py, socket)?;
36 sendto_byte(py, socket, address)
37}
38
39static SELECTORS_EVENT_READ: GILOnceCell<u8> = GILOnceCell::new();
40
41#[pyclass(module = "rlgym_learn_backend", unsendable)]
42pub struct EnvProcessInterface {
43 agent_id_serde_option: Option<PythonSerde>,
44 action_serde_option: Option<PythonSerde>,
45 obs_serde_option: Option<PythonSerde>,
46 reward_serde_option: Option<PythonSerde>,
47 obs_space_serde_option: Option<PythonSerde>,
48 action_space_serde_option: Option<PythonSerde>,
49 state_serde_option: Option<PythonSerde>,
50 state_metrics_serde_option: Option<PythonSerde>,
51 recalculate_agent_id_every_step: bool,
52 flinks_folder: String,
53 proc_packages: Vec<(PyObject, Shmem, String)>,
54 min_process_steps_per_inference: usize,
55 send_state_to_agent_controllers: bool,
56 should_collect_state_metrics: bool,
57 selector: PyObject,
58 timestep_class: PyObject,
59 proc_id_pid_idx_map: HashMap<String, usize>,
60 pid_idx_current_env_action_list: Vec<Option<EnvAction>>,
61 pid_idx_current_agent_id_list: Vec<Option<Vec<PyObject>>>,
62 pid_idx_prev_timestep_id_list: Vec<Vec<Option<u128>>>,
63 pid_idx_current_obs_list: Vec<Vec<PyObject>>,
64 pid_idx_current_action_list: Vec<Vec<PyObject>>,
65 pid_idx_current_aald_list: Vec<Option<PyObject>>,
66 added_process_obs_data_kv_list: Vec<(Py<PyAny>, (Vec<PyObject>, Vec<PyObject>))>,
67 added_process_state_info_kv_list: Vec<(
68 Py<PyAny>,
69 (Option<Py<PyAny>>, Option<Py<PyDict>>, Option<Py<PyDict>>),
70 )>,
71}
72
73impl EnvProcessInterface {
74 fn get_initial_obs_data_proc<'py>(
75 &mut self,
76 py: Python<'py>,
77 pid_idx: usize,
78 ) -> PyResult<(
79 (PyObject, (Vec<PyObject>, Vec<PyObject>)),
80 (
81 PyObject,
82 (Option<PyObject>, Option<Py<PyDict>>, Option<Py<PyDict>>),
83 ),
84 )> {
85 let mut agent_id_serde_option = self
86 .agent_id_serde_option
87 .take()
88 .map(|serde| serde.into_bound(py));
89 let mut obs_serde_option = self
90 .obs_serde_option
91 .take()
92 .map(|serde| serde.into_bound(py));
93
94 let (parent_end, shmem, proc_id) = self.proc_packages.get(pid_idx).unwrap();
95 let shm_slice = unsafe { &shmem.as_slice()[Event::size_of(None)..] };
96 recvfrom_byte(py, parent_end)?;
97 let mut offset = 0;
98 let n_agents;
99 (n_agents, offset) = retrieve_usize(shm_slice, offset)?;
100 let mut agent_id_list: Vec<PyObject> = Vec::with_capacity(n_agents);
101 let mut obs_list: Vec<PyObject> = Vec::with_capacity(n_agents);
102 let mut agent_id;
103 let mut obs;
104 for _ in 0..n_agents {
105 (agent_id, offset) =
106 retrieve_python(py, shm_slice, offset, &mut agent_id_serde_option)?;
107 agent_id_list.push(agent_id.unbind());
108 (obs, offset) = retrieve_python(py, shm_slice, offset, &mut obs_serde_option)?;
109 obs_list.push(obs.unbind());
110 }
111
112 let state_option;
113 if self.send_state_to_agent_controllers {
114 let mut state_serde_option = self
115 .state_serde_option
116 .take()
117 .map(|serde| serde.into_bound(py));
118 let state;
119 (state, _) = retrieve_python(py, shm_slice, offset, &mut state_serde_option)?;
120 state_option = Some(state.unbind());
121 self.state_serde_option = state_serde_option.map(|serde| serde.unbind());
122 } else {
123 state_option = None;
124 }
125
126 let py_proc_id = proc_id.into_py_any(py)?;
127 self.agent_id_serde_option = agent_id_serde_option.map(|serde| serde.unbind());
128 self.obs_serde_option = obs_serde_option.map(|serde| serde.unbind());
129 Ok((
130 (py_proc_id.clone_ref(py), (agent_id_list, obs_list)),
131 (py_proc_id, (state_option, None, None)),
132 ))
133 }
134
135 fn update_with_initial_obs<'py>(
136 &mut self,
137 py: Python<'py>,
138 ) -> PyResult<(Py<PyDict>, Py<PyDict>)> {
139 let n_procs = self.proc_packages.len();
140 let mut obs_data_kv_list = Vec::with_capacity(n_procs);
141 let mut state_info_kv_list = Vec::with_capacity(n_procs);
142 for pid_idx in 0..n_procs {
143 let ((py_proc_id, (agent_id_list, obs_list)), state_info_kv) =
144 self.get_initial_obs_data_proc(py, pid_idx)?;
145 let n_agents = agent_id_list.len();
146 self.pid_idx_current_agent_id_list
147 .push(Some(clone_list(py, &agent_id_list)));
148 self.pid_idx_current_obs_list
149 .push(clone_list(py, &obs_list));
150 self.pid_idx_prev_timestep_id_list
151 .push(vec![None; n_agents]);
152 obs_data_kv_list.push((py_proc_id, (agent_id_list, obs_list)));
153 state_info_kv_list.push(state_info_kv);
154 }
155 Ok((
156 PyDict::from_sequence(&obs_data_kv_list.into_pyobject(py)?)?.unbind(),
157 PyDict::from_sequence(&state_info_kv_list.into_pyobject(py)?)?.unbind(),
158 ))
159 }
160
161 fn get_space_types<'py>(&mut self, py: Python<'py>) -> PyResult<(PyObject, PyObject)> {
162 let mut obs_space_serde_option = self
163 .obs_space_serde_option
164 .take()
165 .map(|serde| serde.into_bound(py));
166 let mut action_space_serde_option = self
167 .action_space_serde_option
168 .take()
169 .map(|serde| serde.into_bound(py));
170
171 let (parent_end, shmem, _) = self.proc_packages.get_mut(0).unwrap();
172 let (ep_evt, used_bytes) = unsafe {
173 Event::from_existing(shmem.as_ptr()).map_err(|err| {
174 InvalidStateError::new_err(format!("Failed to get event: {}", err.to_string()))
175 })?
176 };
177 let shm_slice = unsafe { &mut shmem.as_slice_mut()[used_bytes..] };
178 append_header(shm_slice, 0, Header::EnvShapesRequest);
179 ep_evt
180 .set(EventState::Signaled)
181 .map_err(|err| InvalidStateError::new_err(err.to_string()))?;
182 recvfrom_byte(py, parent_end)?;
183 let mut offset = 0;
184 let obs_space;
185 (obs_space, offset) = retrieve_python(py, shm_slice, offset, &mut obs_space_serde_option)?;
186 let action_space;
187 (action_space, _) = retrieve_python(py, shm_slice, offset, &mut action_space_serde_option)?;
188
189 self.obs_space_serde_option = obs_space_serde_option.map(|serde| serde.unbind());
190 self.action_space_serde_option = action_space_serde_option.map(|serde| serde.unbind());
191 Ok((obs_space.unbind(), action_space.unbind()))
192 }
193
194 fn add_proc_package<'py>(
195 &mut self,
196 py: Python<'py>,
197 proc_package_def: (PyObject, PyObject, PyObject, String),
198 ) -> PyResult<()> {
199 let (_, parent_end, child_sockname, proc_id) = proc_package_def;
200 sync_with_env_process(py, &parent_end, &child_sockname)?;
201 let flink = get_flink(&self.flinks_folder[..], proc_id.as_str());
202 let shmem = ShmemConf::new()
203 .flink(flink.clone())
204 .open()
205 .map_err(|err| {
206 InvalidStateError::new_err(format!("Unable to open shmem flink {}: {}", flink, err))
207 })?;
208 self.selector.call_method1(
209 py,
210 intern!(py, "register"),
211 (
212 parent_end.clone_ref(py),
213 SELECTORS_EVENT_READ.get_or_init(py, || {
214 PyModule::import(py, "selectors")
215 .unwrap()
216 .getattr("EVENT_READ")
217 .unwrap()
218 .extract()
219 .unwrap()
220 }),
221 self.proc_packages.len(),
222 ),
223 )?;
224 self.proc_id_pid_idx_map
225 .insert(proc_id.clone(), self.proc_packages.len());
226 self.proc_packages.push((parent_end, shmem, proc_id));
227
228 Ok(())
229 }
230
231 fn collect_response(
236 &mut self,
237 pid_idx: usize,
238 ) -> PyResult<(
239 usize,
240 (PyObject, (Vec<PyObject>, Vec<PyObject>)),
241 (
242 PyObject,
243 (Vec<PyObject>, PyObject, Option<PyObject>, Option<PyObject>),
244 ),
245 (
246 PyObject,
247 (Option<PyObject>, Option<Py<PyDict>>, Option<Py<PyDict>>),
248 ),
249 )> {
250 let env_action = self.pid_idx_current_env_action_list[pid_idx]
251 .as_ref()
252 .ok_or_else(|| {
253 InvalidStateError::new_err(
254 "Tried to collect response from env which doesn't have an env action yet",
255 )
256 })?;
257 let is_step_action = matches!(env_action, EnvAction::STEP { .. });
258 let new_episode = !is_step_action;
259 let (_, shmem, proc_id) = self.proc_packages.get(pid_idx).unwrap();
260 let evt_used_bytes = Event::size_of(None);
261 let shm_slice = unsafe { &shmem.as_slice()[evt_used_bytes..] };
262 let mut offset = 0;
263 Python::with_gil(|py| {
264 let current_agent_id_list = self
265 .pid_idx_current_agent_id_list
266 .get_mut(pid_idx)
267 .unwrap()
268 .take()
269 .unwrap();
270
271 let mut agent_id_serde_option = self
272 .agent_id_serde_option
273 .take()
274 .map(|serde| serde.into_bound(py));
275 let mut obs_serde_option = self
276 .obs_serde_option
277 .take()
278 .map(|serde| serde.into_bound(py));
279 let mut reward_serde_option = self
280 .reward_serde_option
281 .take()
282 .map(|serde| serde.into_bound(py));
283
284 let n_agents;
286 let (
287 mut agent_id_list,
288 mut obs_list,
289 mut reward_list_option,
290 mut terminated_list_option,
291 mut truncated_list_option,
292 );
293
294 if new_episode {
295 (n_agents, offset) = retrieve_usize(shm_slice, offset)?;
296 agent_id_list = Vec::with_capacity(n_agents);
297 } else {
298 n_agents = current_agent_id_list.len();
299 if self.recalculate_agent_id_every_step {
300 agent_id_list = Vec::with_capacity(n_agents);
301 } else {
302 agent_id_list = current_agent_id_list;
303 }
304 }
305 obs_list = Vec::with_capacity(n_agents);
306 if is_step_action {
307 reward_list_option = Some(Vec::with_capacity(n_agents));
308 terminated_list_option = Some(Vec::with_capacity(n_agents));
309 truncated_list_option = Some(Vec::with_capacity(n_agents));
310 } else {
311 reward_list_option = None;
312 terminated_list_option = None;
313 truncated_list_option = None;
314 }
315
316 for _ in 0..n_agents {
318 if self.recalculate_agent_id_every_step || new_episode {
319 let agent_id;
320 (agent_id, offset) =
321 retrieve_python(py, shm_slice, offset, &mut agent_id_serde_option)?;
322 agent_id_list.push(agent_id.unbind());
323 }
324 let obs;
325 (obs, offset) = retrieve_python(py, shm_slice, offset, &mut obs_serde_option)?;
326 obs_list.push(obs.unbind());
327 if is_step_action {
328 let reward;
329 (reward, offset) =
330 retrieve_python(py, shm_slice, offset, &mut reward_serde_option)?;
331 reward_list_option.as_mut().unwrap().push(reward.unbind());
332 let terminated;
333 (terminated, offset) = retrieve_bool(shm_slice, offset)?;
334 terminated_list_option.as_mut().unwrap().push(terminated);
335 let truncated;
336 (truncated, offset) = retrieve_bool(shm_slice, offset)?;
337 truncated_list_option.as_mut().unwrap().push(truncated);
338 }
339 }
340
341 let state_option;
342 if self.send_state_to_agent_controllers {
343 let mut state_serde_option = self
344 .state_serde_option
345 .take()
346 .map(|serde| serde.into_bound(py));
347 let state;
348 (state, offset) = retrieve_python(py, shm_slice, offset, &mut state_serde_option)?;
349 state_option = Some(state.unbind());
350 self.state_serde_option = state_serde_option.map(|serde| serde.unbind());
351 } else {
352 state_option = None;
353 }
354
355 let metrics_option;
356 if self.should_collect_state_metrics {
357 let mut state_metrics_serde_option = self
358 .state_metrics_serde_option
359 .take()
360 .map(|serde| serde.into_bound(py));
361 let state_metrics;
362 (state_metrics, offset) =
363 retrieve_python(py, shm_slice, offset, &mut state_metrics_serde_option)?;
364 metrics_option = Some(state_metrics.unbind());
365 self.state_metrics_serde_option =
366 state_metrics_serde_option.map(|serde| serde.unbind());
367 } else {
368 metrics_option = None;
369 }
370
371 let timestep_id_list_option;
372 let mut timestep_list;
373 if is_step_action {
374 let timestep_class = self.timestep_class.bind(py);
375 let mut timestep_id_list = Vec::with_capacity(n_agents);
376 timestep_list = Vec::with_capacity(n_agents);
377 for (
378 prev_timestep_id,
379 agent_id,
380 obs,
381 next_obs,
382 action,
383 reward,
384 &terminated,
385 &truncated,
386 ) in izip!(
387 self.pid_idx_prev_timestep_id_list.get(pid_idx).unwrap(),
388 &agent_id_list,
389 self.pid_idx_current_obs_list.get(pid_idx).unwrap(),
390 &obs_list,
391 &self.pid_idx_current_action_list[pid_idx],
392 reward_list_option.as_ref().unwrap(),
393 terminated_list_option.as_ref().unwrap(),
394 truncated_list_option.as_ref().unwrap()
395 ) {
396 let timestep_id = fastrand::u128(..);
397 timestep_id_list.push(Some(timestep_id));
398 timestep_list.push(
399 timestep_class
400 .call1((
401 proc_id.into_py_any(py)?,
402 timestep_id,
403 *prev_timestep_id,
404 agent_id.clone_ref(py),
405 obs,
406 next_obs,
407 action,
408 reward,
409 terminated,
410 truncated,
411 ))?
412 .unbind(),
413 );
414 }
415 timestep_id_list_option = Some(timestep_id_list);
416 } else {
417 timestep_id_list_option = None;
418 timestep_list = Vec::new();
419 }
420 let n_timesteps = timestep_list.len();
421
422 let terminated_dict_option;
423 let truncated_dict_option;
424 if new_episode {
425 terminated_dict_option = None;
426 truncated_dict_option = None;
427 } else {
428 let mut terminated_kv_list = Vec::with_capacity(n_agents);
429 let mut truncated_kv_list = Vec::with_capacity(n_agents);
430 for (agent_id, terminated, truncated) in izip!(
431 &agent_id_list,
432 terminated_list_option.unwrap(),
433 truncated_list_option.unwrap()
434 ) {
435 terminated_kv_list.push((agent_id.clone_ref(py), terminated));
436 truncated_kv_list.push((agent_id.clone_ref(py), truncated));
437 }
438 terminated_dict_option =
439 Some(PyDict::from_sequence(&terminated_kv_list.into_pyobject(py)?)?.unbind());
440 truncated_dict_option =
441 Some(PyDict::from_sequence(&truncated_kv_list.into_pyobject(py)?)?.unbind());
442 }
443
444 let prev_timestep_id_list = &mut self.pid_idx_prev_timestep_id_list[pid_idx];
446 if is_step_action {
447 prev_timestep_id_list.clear();
448 prev_timestep_id_list.append(&mut timestep_id_list_option.unwrap());
449 } else if let EnvAction::SET_STATE {
450 prev_timestep_id_dict_option: Some(prev_timestep_id_dict),
451 ..
452 } = env_action
453 {
454 let prev_timestep_id_dict = prev_timestep_id_dict.downcast_bound::<PyDict>(py)?;
455 prev_timestep_id_list.clear();
456 for agent_id in agent_id_list.iter() {
457 let agent_id = agent_id.bind(py);
458 prev_timestep_id_list.push(
459 prev_timestep_id_dict
460 .get_item(agent_id)?
461 .map_or(Ok(None), |prev_timestep_id| {
462 prev_timestep_id.extract::<Option<u128>>()
463 })?,
464 );
465 }
466 } else {
467 prev_timestep_id_list.clear();
468 prev_timestep_id_list.append(&mut vec![None; n_agents]);
469 }
470 self.pid_idx_current_agent_id_list[pid_idx] = Some(clone_list(py, &agent_id_list));
471 self.pid_idx_current_obs_list[pid_idx] = clone_list(py, &obs_list);
472
473 let py_proc_id = proc_id.into_py_any(py)?;
474 let obs_data_kv = (py_proc_id.clone_ref(py), (agent_id_list, obs_list));
475 let timestep_data_kv = (
476 py_proc_id.clone_ref(py),
477 (
478 timestep_list,
479 (&self.pid_idx_current_aald_list[pid_idx]).into_py_any(py)?,
480 metrics_option,
481 state_option.as_ref().map(|state| state.clone_ref(py)),
482 ),
483 );
484 let state_info_kv = (
485 py_proc_id,
486 (state_option, terminated_dict_option, truncated_dict_option),
487 );
488
489 self.agent_id_serde_option = agent_id_serde_option.map(|serde| serde.unbind());
490 self.obs_serde_option = obs_serde_option.map(|serde| serde.unbind());
491 self.reward_serde_option = reward_serde_option.map(|serde| serde.unbind());
492
493 Ok((n_timesteps, obs_data_kv, timestep_data_kv, state_info_kv))
494 })
495 }
496}
497
498#[pymethods]
499impl EnvProcessInterface {
500 #[new]
501 #[pyo3(signature = (
502 agent_id_serde_option,
503 action_serde_option,
504 obs_serde_option,
505 reward_serde_option,
506 obs_space_serde_option,
507 action_space_serde_option,
508 state_serde_option,
509 state_metrics_serde_option,
510 recalculate_agent_id_every_step,
511 flinks_folder,
512 min_process_steps_per_inference,
513 send_state_to_agent_controllers,
514 should_collect_state_metrics,
515 ))]
516 fn new(
517 agent_id_serde_option: Option<PythonSerde>,
518 action_serde_option: Option<PythonSerde>,
519 obs_serde_option: Option<PythonSerde>,
520 reward_serde_option: Option<PythonSerde>,
521 obs_space_serde_option: Option<PythonSerde>,
522 action_space_serde_option: Option<PythonSerde>,
523 state_serde_option: Option<PythonSerde>,
524 state_metrics_serde_option: Option<PythonSerde>,
525 recalculate_agent_id_every_step: bool,
526 flinks_folder: String,
527 min_process_steps_per_inference: usize,
528 send_state_to_agent_controllers: bool,
529 should_collect_state_metrics: bool,
530 ) -> PyResult<Self> {
531 Python::with_gil::<_, PyResult<Self>>(|py| {
532 let timestep_class = PyModule::import(py, "rlgym_learn.experience.timestep")?
533 .getattr("Timestep")?
534 .unbind();
535 let selector = PyModule::import(py, "selectors")?
536 .getattr("DefaultSelector")?
537 .call0()?
538 .unbind();
539 Ok(EnvProcessInterface {
540 agent_id_serde_option,
541 action_serde_option,
542 obs_serde_option,
543 reward_serde_option,
544 obs_space_serde_option,
545 action_space_serde_option,
546 state_serde_option,
547 state_metrics_serde_option,
548 recalculate_agent_id_every_step,
549 flinks_folder,
550 proc_packages: Vec::new(),
551 min_process_steps_per_inference,
552 send_state_to_agent_controllers,
553 should_collect_state_metrics,
554 selector,
555 timestep_class,
556 proc_id_pid_idx_map: HashMap::new(),
557 pid_idx_current_env_action_list: Vec::new(),
558 pid_idx_current_agent_id_list: Vec::new(),
559 pid_idx_prev_timestep_id_list: Vec::new(),
560 pid_idx_current_obs_list: Vec::new(),
561 pid_idx_current_action_list: Vec::new(),
562 pid_idx_current_aald_list: Vec::new(),
563 added_process_obs_data_kv_list: Vec::new(),
564 added_process_state_info_kv_list: Vec::new(),
565 })
566 })
567 }
568
569 fn init_processes(
578 &mut self,
579 proc_package_defs: Vec<(PyObject, PyObject, PyObject, String)>,
580 ) -> PyResult<(Py<PyDict>, Py<PyDict>, PyObject, PyObject)> {
581 Python::with_gil(|py| {
582 proc_package_defs
583 .into_iter()
584 .try_for_each::<_, PyResult<()>>(|proc_package_def| {
585 self.add_proc_package(py, proc_package_def)
586 })?;
587 let (initial_obs_data_dict, initial_state_info_dict) =
588 self.update_with_initial_obs(py)?;
589 let n_procs = self.proc_packages.len();
590 self.min_process_steps_per_inference =
591 min(self.min_process_steps_per_inference, n_procs);
592 for _ in 0..n_procs {
593 self.pid_idx_current_env_action_list.push(None);
594 self.pid_idx_current_action_list.push(Vec::new());
595 self.pid_idx_current_aald_list.push(None);
596 }
597 let (obs_space, action_space) = self.get_space_types(py)?;
598
599 Ok((
600 initial_obs_data_dict,
601 initial_state_info_dict,
602 obs_space,
603 action_space,
604 ))
605 })
606 }
607
608 fn add_process(
609 &mut self,
610 proc_package_def: (PyObject, PyObject, PyObject, String),
611 ) -> PyResult<()> {
612 Python::with_gil(|py| {
613 let pid_idx = self.proc_packages.len();
614 self.add_proc_package(py, proc_package_def)?;
615 let ((py_proc_id, (agent_id_list, obs_list)), state_info_kv) =
616 self.get_initial_obs_data_proc(py, pid_idx)?;
617 let n_agents = agent_id_list.len();
618 self.pid_idx_current_agent_id_list
619 .push(Some(clone_list(py, &agent_id_list)));
620 self.pid_idx_current_obs_list
621 .push(clone_list(py, &obs_list));
622 self.pid_idx_prev_timestep_id_list
623 .push(vec![None; n_agents]);
624 self.pid_idx_current_env_action_list.push(None);
625 self.pid_idx_current_action_list
626 .push(Vec::with_capacity(n_agents));
627 self.pid_idx_current_aald_list.push(None);
628 self.added_process_obs_data_kv_list
629 .push((py_proc_id, (agent_id_list, obs_list)));
630 self.added_process_state_info_kv_list.push(state_info_kv);
631 Ok(())
632 })
633 }
634
635 fn delete_process(&mut self) -> PyResult<()> {
636 let (parent_end, mut shmem, proc_id) = self.proc_packages.pop().unwrap();
637 self.proc_id_pid_idx_map.remove(&proc_id);
638 let (ep_evt, used_bytes) = unsafe {
639 Event::from_existing(shmem.as_ptr()).map_err(|err| {
640 InvalidStateError::new_err(format!("Failed to get event: {}", err.to_string()))
641 })?
642 };
643 let shm_slice = unsafe { &mut shmem.as_slice_mut()[used_bytes..] };
644 append_header(shm_slice, 0, Header::Stop);
645 ep_evt
646 .set(EventState::Signaled)
647 .map_err(|err| InvalidStateError::new_err(err.to_string()))?;
648 self.pid_idx_current_agent_id_list.pop();
649 self.pid_idx_prev_timestep_id_list.pop();
650 self.pid_idx_current_obs_list.pop();
651 self.pid_idx_current_env_action_list.pop();
652 self.pid_idx_current_action_list.pop();
653 self.pid_idx_current_aald_list.pop();
654 self.added_process_state_info_kv_list
655 .retain(|(py_proc_id, _)| py_proc_id.to_string() != proc_id);
656 self.min_process_steps_per_inference = min(
657 self.min_process_steps_per_inference,
658 self.proc_packages.len().try_into().unwrap(),
659 );
660 Python::with_gil(|py| {
661 self.selector
662 .call_method1(py, intern!(py, "unregister"), (parent_end,))?;
663 Ok(())
664 })
665 }
666
667 fn increase_min_process_steps_per_inference(&mut self) -> usize {
668 self.min_process_steps_per_inference = min(
669 self.min_process_steps_per_inference + 1,
670 self.proc_packages.len().try_into().unwrap(),
671 );
672 self.min_process_steps_per_inference
673 }
674
675 fn decrease_min_process_steps_per_inference(&mut self) -> usize {
676 self.min_process_steps_per_inference = max(self.min_process_steps_per_inference - 1, 1);
677 self.min_process_steps_per_inference
678 }
679
680 fn cleanup(&mut self) -> PyResult<()> {
681 while let Some(proc_package) = self.proc_packages.pop() {
682 let (parent_end, mut shmem, _) = proc_package;
683 let (ep_evt, used_bytes) = unsafe {
684 Event::from_existing(shmem.as_ptr()).map_err(|err| {
685 InvalidStateError::new_err(format!("Failed to get event: {}", err.to_string()))
686 })?
687 };
688 let shm_slice = unsafe { &mut shmem.as_slice_mut()[used_bytes..] };
689 append_header(shm_slice, 0, Header::Stop);
690 ep_evt
691 .set(EventState::Signaled)
692 .map_err(|err| InvalidStateError::new_err(err.to_string()))?;
693 Python::with_gil(|py| {
694 self.selector
695 .call_method1(py, intern!(py, "unregister"), (parent_end,))
696 })?;
697 thread::sleep(Duration::from_millis(1));
699 }
700 self.proc_id_pid_idx_map.clear();
701 self.pid_idx_current_agent_id_list.clear();
702 self.pid_idx_prev_timestep_id_list.clear();
703 self.pid_idx_current_obs_list.clear();
704 self.pid_idx_current_action_list.clear();
705 self.pid_idx_current_aald_list.clear();
706 self.added_process_state_info_kv_list.clear();
707 Ok(())
708 }
709
710 fn collect_step_data(&mut self) -> PyResult<(usize, Py<PyDict>, Py<PyDict>, Py<PyDict>)> {
717 let mut n_process_steps_collected = 0;
718 let mut total_timesteps_collected = 0;
719 let mut obs_data_kv_list = Vec::with_capacity(self.min_process_steps_per_inference);
720 let mut timestep_data_kv_list = Vec::with_capacity(self.min_process_steps_per_inference);
721 let mut state_info_kv_list = Vec::with_capacity(
722 self.min_process_steps_per_inference + self.added_process_state_info_kv_list.len(),
723 );
724 obs_data_kv_list.append(&mut self.added_process_obs_data_kv_list);
725 state_info_kv_list.append(&mut self.added_process_state_info_kv_list);
726 Python::with_gil(|py| {
727 while n_process_steps_collected < self.min_process_steps_per_inference {
728 for (key, event) in self
729 .selector
730 .bind(py)
731 .call_method0(intern!(py, "select"))?
732 .extract::<Vec<(PyObject, u8)>>()?
733 {
734 if event & SELECTORS_EVENT_READ.get(py).unwrap() == 0 {
735 continue;
736 }
737 let (parent_end, _, _, pid_idx) =
738 key.extract::<(PyObject, PyObject, PyObject, usize)>(py)?;
739 recvfrom_byte(py, &parent_end)?;
740 let (n_timesteps, obs_data_kv, timestep_data_kv, state_info_kv) =
741 self.collect_response(pid_idx)?;
742 obs_data_kv_list.push(obs_data_kv);
743 timestep_data_kv_list.push(timestep_data_kv);
744 state_info_kv_list.push(state_info_kv);
745 n_process_steps_collected += 1;
746 total_timesteps_collected += n_timesteps;
747 }
748 }
749 Ok((
750 total_timesteps_collected,
751 PyDict::from_sequence(&obs_data_kv_list.into_pyobject(py)?)?.unbind(),
752 PyDict::from_sequence(×tep_data_kv_list.into_pyobject(py)?)?.unbind(),
753 PyDict::from_sequence(&state_info_kv_list.into_pyobject(py)?)?.unbind(),
754 ))
755 })
756 }
757
758 fn send_env_actions(&mut self, env_actions: HashMap<String, EnvAction>) -> PyResult<()> {
759 Python::with_gil(|py| {
760 let mut action_serde_option = self
761 .action_serde_option
762 .take()
763 .map(|serde| serde.into_bound(py));
764 let mut state_serde_option = self
765 .state_serde_option
766 .take()
767 .map(|serde| serde.into_bound(py));
768
769 for (proc_id, env_action) in env_actions.into_iter() {
770 let &pid_idx = self.proc_id_pid_idx_map.get(&proc_id).unwrap();
771 let (_, shmem, _) = self.proc_packages.get_mut(pid_idx).unwrap();
772 let (ep_evt, evt_used_bytes) = unsafe {
773 Event::from_existing(shmem.as_ptr()).map_err(|err| {
774 InvalidStateError::new_err(format!(
775 "Failed to get event from epi to process with index {}: {}",
776 pid_idx,
777 err.to_string()
778 ))
779 })?
780 };
781 let shm_slice = unsafe { &mut shmem.as_slice_mut()[evt_used_bytes..] };
782
783 if let EnvAction::STEP {
784 ref action_list,
785 ref action_associated_learning_data,
786 } = env_action
787 {
788 let current_action_list = &mut self.pid_idx_current_action_list[pid_idx];
789 current_action_list.clear();
790 current_action_list.append(
791 &mut action_list
792 .bind(py)
793 .iter()
794 .map(|action| action.unbind())
795 .collect_vec(),
796 );
797 self.pid_idx_current_aald_list[pid_idx] =
798 Some(action_associated_learning_data.clone_ref(py));
799 } else {
800 self.pid_idx_current_aald_list[pid_idx] = None;
801 }
802
803 let offset = append_header(shm_slice, 0, Header::EnvAction);
804 _ = append_env_action_new(
805 py,
806 shm_slice,
807 offset,
808 &env_action,
809 &mut action_serde_option,
810 &mut state_serde_option,
811 )?;
812
813 ep_evt
814 .set(EventState::Signaled)
815 .map_err(|err| InvalidStateError::new_err(err.to_string()))?;
816 self.pid_idx_current_env_action_list[pid_idx] = Some(env_action);
817 }
818 Ok(())
819 })
820 }
821}