rlgym_learn_backend/
communication.rs

1use std::fmt::{self, Display, Formatter};
2use std::mem::size_of;
3use std::os::raw::c_double;
4
5use pyo3::exceptions::asyncio::InvalidStateError;
6use pyo3::sync::GILOnceCell;
7use pyo3::types::PyBytes;
8use pyo3::{intern, prelude::*, IntoPyObjectExt};
9
10use paste::paste;
11
12use crate::serdes::pyany_serde::{detect_pyany_serde, get_pyany_serde, BoundPythonSerde};
13use crate::serdes::serde_enum::retrieve_serde;
14
15#[derive(Debug, PartialEq)]
16pub enum Header {
17    EnvShapesRequest,
18    EnvAction,
19    Stop,
20}
21
22impl Display for Header {
23    fn fmt(&self, f: &mut Formatter<'_>) -> fmt::Result {
24        match self {
25            Self::EnvShapesRequest => write!(f, "EnvShapesRequest"),
26            Self::EnvAction => write!(f, "EnvAction"),
27            Self::Stop => write!(f, "Stop"),
28        }
29    }
30}
31
32static INTERNED_INT_1: GILOnceCell<PyObject> = GILOnceCell::new();
33static INTERNED_BYTES_0: GILOnceCell<PyObject> = GILOnceCell::new();
34
35#[pyfunction]
36pub fn recvfrom_byte_py(socket: PyObject) -> PyResult<PyObject> {
37    Python::with_gil(|py| recvfrom_byte(py, &socket))
38}
39
40pub fn recvfrom_byte<'py>(py: Python<'py>, socket: &PyObject) -> PyResult<PyObject> {
41    socket.call_method1(
42        py,
43        intern!(py, "recvfrom"),
44        (INTERNED_INT_1.get_or_init(py, || 1_i64.into_py_any(py).unwrap()),),
45    )
46}
47
48#[pyfunction]
49pub fn sendto_byte_py(socket: PyObject, address: PyObject) -> PyResult<()> {
50    Python::with_gil(|py| sendto_byte(py, &socket, &address))
51}
52
53pub fn sendto_byte<'py>(py: Python<'py>, socket: &PyObject, address: &PyObject) -> PyResult<()> {
54    socket.call_method1(
55        py,
56        intern!(py, "sendto"),
57        (
58            INTERNED_BYTES_0
59                .get_or_init(py, || PyBytes::new(py, &vec![0_u8][..]).into_any().unbind()),
60            address,
61        ),
62    )?;
63    Ok(())
64}
65
66pub fn get_flink(flinks_folder: &str, proc_id: &str) -> String {
67    format!("{}/{}", flinks_folder, proc_id)
68}
69
70pub fn append_header(buf: &mut [u8], offset: usize, header: Header) -> usize {
71    buf[offset] = match header {
72        Header::EnvShapesRequest => 0,
73        Header::EnvAction => 1,
74        Header::Stop => 2,
75    };
76    offset + 1
77}
78
79pub fn retrieve_header(slice: &[u8], offset: usize) -> PyResult<(Header, usize)> {
80    let header = match slice[offset] {
81        0 => Ok(Header::EnvShapesRequest),
82        1 => Ok(Header::EnvAction),
83        2 => Ok(Header::Stop),
84        v => Err(InvalidStateError::new_err(format!(
85            "tried to retrieve header from shared_memory but got value {}",
86            v
87        ))),
88    }?;
89    Ok((header, offset + 1))
90}
91
92macro_rules! define_primitive_communication {
93    ($type:ty) => {
94        paste! {
95            pub fn [<append_ $type>](buf: &mut [u8], offset: usize, val: $type) -> usize {
96                let end = offset + size_of::<$type>();
97                buf[offset..end].copy_from_slice(&val.to_ne_bytes());
98                end
99            }
100
101            pub fn [<retrieve_ $type>](buf: &[u8], offset: usize) -> PyResult<($type, usize)> {
102                let end = offset + size_of::<$type>();
103                Ok(($type::from_ne_bytes(buf[offset..end].try_into()?), end))
104            }
105        }
106    };
107}
108
109define_primitive_communication!(usize);
110define_primitive_communication!(c_double);
111define_primitive_communication!(i64);
112define_primitive_communication!(u64);
113define_primitive_communication!(f32);
114define_primitive_communication!(f64);
115
116pub fn append_bool(buf: &mut [u8], offset: usize, val: bool) -> usize {
117    let end = offset + size_of::<u8>();
118    let u8_bool = if val { 1_u8 } else { 0 };
119    buf[offset..end].copy_from_slice(&u8_bool.to_ne_bytes());
120    end
121}
122
123pub fn retrieve_bool(slice: &[u8], offset: usize) -> PyResult<(bool, usize)> {
124    let end = offset + size_of::<bool>();
125    let val = match u8::from_ne_bytes(slice[offset..end].try_into()?) {
126        0 => Ok(false),
127        1 => Ok(true),
128        v => Err(InvalidStateError::new_err(format!(
129            "tried to retrieve bool from shared_memory but got value {}",
130            v
131        ))),
132    }?;
133    Ok((val, end))
134}
135
136#[macro_export]
137macro_rules! append_n_vec_elements {
138    ($buf: ident, $offset: expr, $vec: ident, $n: expr) => {{
139        let mut offset = $offset;
140        for idx in 0..$n {
141            offset = crate::communication::append_f32($buf, offset, $vec[idx]);
142        }
143        offset
144    }};
145}
146
147#[macro_export]
148macro_rules! retrieve_n_vec_elements {
149    ($buf: ident, $offset: expr, $n: expr) => {{
150        let mut offset = $offset;
151        let mut val;
152        let mut vec = Vec::with_capacity($n);
153        for _ in 0..$n {
154            (val, offset) = crate::communication::retrieve_f32($buf, offset).unwrap();
155            vec.push(val);
156        }
157        (vec, offset)
158    }};
159}
160
161#[macro_export]
162macro_rules! append_n_vec_elements_option {
163    ($buf: ident, $offset: expr, $vec_option: ident, $n: expr) => {{
164        let mut offset = $offset;
165        if let Some(vec) = $vec_option {
166            offset = crate::communication::append_bool($buf, offset, true);
167            for idx in 0..$n {
168                offset = crate::communication::append_f32($buf, offset, vec[idx]);
169            }
170        } else {
171            offset = crate::communication::append_bool($buf, offset, false)
172        }
173        offset
174    }};
175}
176
177#[macro_export]
178macro_rules! retrieve_n_vec_elements_option {
179    ($buf: ident, $offset: expr, $n: expr) => {{
180        let mut offset = $offset;
181        let is_some;
182        (is_some, offset) = crate::communication::retrieve_bool($buf, offset).unwrap();
183        if is_some {
184            let mut val;
185            let mut vec = Vec::with_capacity($n);
186            for _ in 0..$n {
187                (val, offset) = crate::communication::retrieve_f32($buf, offset).unwrap();
188                vec.push(val);
189            }
190            (Some(vec), offset)
191        } else {
192            (None, offset)
193        }
194    }};
195}
196
197pub fn insert_bytes(buf: &mut [u8], offset: usize, bytes: &[u8]) -> PyResult<usize> {
198    let end = offset + bytes.len();
199    buf[offset..end].copy_from_slice(bytes);
200    Ok(end)
201}
202
203pub fn append_bytes(buf: &mut [u8], offset: usize, bytes: &[u8]) -> PyResult<usize> {
204    let bytes_len = bytes.len();
205    let start = append_usize(buf, offset, bytes_len);
206    let end = start + bytes.len();
207    buf[start..end].copy_from_slice(bytes);
208    Ok(end)
209}
210
211pub fn retrieve_bytes(slice: &[u8], offset: usize) -> PyResult<(&[u8], usize)> {
212    let (len, start) = retrieve_usize(slice, offset)?;
213    let end = start + len;
214    Ok((&slice[start..end], end))
215}
216
217pub fn append_python<'py1, 'py2>(
218    buf: &mut [u8],
219    offset: usize,
220    obj: &Bound<'py1, PyAny>,
221    python_serde_option: &mut Option<BoundPythonSerde<'py2>>,
222) -> PyResult<usize> {
223    let mut offset = offset;
224    match python_serde_option {
225        Some(BoundPythonSerde::TypeSerde(type_serde)) => {
226            offset = append_bytes(
227                buf,
228                offset,
229                type_serde
230                    .call_method1(intern!(obj.py(), "to_bytes"), (obj,))?
231                    .downcast::<PyBytes>()?
232                    .as_bytes(),
233            )?;
234        }
235        Some(BoundPythonSerde::PyAnySerde(pyany_serde)) => {
236            let serde_enum_bytes = pyany_serde.get_enum_bytes();
237            let end = offset + serde_enum_bytes.len();
238            buf[offset..end].copy_from_slice(&serde_enum_bytes[..]);
239            offset = pyany_serde.append(buf, end, &obj)?;
240        }
241        None => {
242            let mut new_pyany_serde = detect_pyany_serde(&obj)?;
243            let serde_enum_bytes = new_pyany_serde.get_enum_bytes();
244            let end = offset + serde_enum_bytes.len();
245            buf[offset..end].copy_from_slice(&serde_enum_bytes[..]);
246            offset = new_pyany_serde.append(buf, end, &obj)?;
247            *python_serde_option = Some(BoundPythonSerde::PyAnySerde(new_pyany_serde));
248        }
249    }
250    return Ok(offset);
251}
252
253pub fn append_python_option<'py>(
254    buf: &mut [u8],
255    offset: usize,
256    obj_option: &Option<&Bound<'py, PyAny>>,
257    python_serde_option: &mut Option<BoundPythonSerde<'py>>,
258) -> PyResult<usize> {
259    let mut offset = offset;
260    if let Some(obj) = obj_option {
261        offset = append_bool(buf, offset, true);
262        offset = append_python(buf, offset, obj, python_serde_option)?;
263    } else {
264        offset = append_bool(buf, offset, false);
265    }
266    Ok(offset)
267}
268
269pub fn retrieve_python<'py1, 'py2: 'py1>(
270    py: Python<'py1>,
271    buf: &[u8],
272    offset: usize,
273    python_serde_option: &mut Option<BoundPythonSerde<'py2>>,
274) -> PyResult<(Bound<'py1, PyAny>, usize)> {
275    let obj;
276    let mut offset = offset;
277    match python_serde_option {
278        Some(BoundPythonSerde::TypeSerde(type_serde)) => {
279            let obj_bytes;
280            (obj_bytes, offset) = retrieve_bytes(buf, offset)?;
281            obj = type_serde
282                .call_method1(intern!(py, "from_bytes"), (PyBytes::new(py, obj_bytes),))?;
283        }
284        Some(BoundPythonSerde::PyAnySerde(pyany_serde)) => {
285            offset += pyany_serde.get_enum_bytes().len();
286            (obj, offset) = pyany_serde.retrieve(py, buf, offset)?;
287        }
288        None => {
289            let serde;
290            (serde, offset) = retrieve_serde(buf, offset)?;
291            let mut new_pyany_serde = get_pyany_serde(serde)?;
292            (obj, offset) = new_pyany_serde.retrieve(py, buf, offset)?;
293            *python_serde_option = Some(BoundPythonSerde::PyAnySerde(new_pyany_serde));
294        }
295    }
296    return Ok((IntoPyObject::into_pyobject(obj, py)?, offset));
297}
298
299pub fn retrieve_python_option<'py1, 'py2: 'py1>(
300    py: Python<'py1>,
301    buf: &[u8],
302    offset: usize,
303    python_serde_option: &mut Option<BoundPythonSerde<'py2>>,
304) -> PyResult<(Option<Bound<'py1, PyAny>>, usize)> {
305    let mut offset = offset;
306    let is_some;
307    (is_some, offset) = retrieve_bool(buf, offset)?;
308    if is_some {
309        let (obj, offset) = retrieve_python(py, buf, offset, python_serde_option)?;
310        Ok((Some(obj), offset))
311    } else {
312        Ok((None, offset))
313    }
314}