rlgym_learn_backend/
env_process_interface.rs

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    // Returns number of timesteps collected, plus three kv pairs: the keys are all the proc id,
232    // and the values are (agent id list, obs list),
233    // (timestep list, optional state metrics, optional state),
234    // and (optional state, optional terminated dict, optional truncated dict) respectively
235    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            // Get n_agents for incoming data and instantiate lists
285            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            // Populate lists
317            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            // Set prev_timestep_id_list for proc
445            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    // Return (
570    // Dict with key being proc_id, and value being (
571    // list of AgentID,
572    // list of ObsType
573    // ),
574    // ObsSpaceType,
575    // ActionSpaceType
576    // )
577    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            // This sleep seems to be needed for the shared memory to get set/read correctly
698            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    // Returns: (
711    // list of AgentID
712    // list of ObsType
713    // Dict of timesteps, state metrics, and state by proc id
714    // Dict of state, terminated dict, and truncated dict by proc id
715    // )
716    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(&timestep_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}