Skip to main content

moirai_transport/
remote_task.rs

1//! Remote task envelopes over transport-owned bytes.
2
3use crate::{
4    Address, NetworkTransport, RemoteAddress, Transport, TransportError, TransportResult,
5    payload::{ServerPayloadRegion, TransportPayload, archive_transport_payload},
6    safe_channel::{ArchiveSerialize, ArchiveView, ArchivedMessage},
7};
8
9mod capability;
10mod server;
11
12pub use capability::{
13    EchoBytesCapability, IntoRemoteOperation, RemoteCapability, RemoteCapabilityToken,
14    RemoteTaskOperationKind, SumU64Capability, build_remote_operation,
15};
16pub use server::{
17    BoundedRemoteTaskServer, RemoteTaskQueueCapacity, RemoteTaskRequestLimit, RemoteTaskServer,
18    RemoteTaskServerStats, RemoteTaskWorkerCount,
19};
20
21const OP_ECHO_BYTES: u8 = 1;
22const OP_SUM_U64: u8 = 2;
23const RESULT_BYTES: u8 = 1;
24const RESULT_U64: u8 = 2;
25
26/// Remote task identifier.
27#[repr(transparent)]
28#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
29pub struct RemoteTaskId(u64);
30
31impl RemoteTaskId {
32    /// Construct a remote task id.
33    pub const fn new(id: u64) -> Self {
34        Self(id)
35    }
36
37    /// Return the raw id.
38    pub const fn get(self) -> u64 {
39        self.0
40    }
41}
42
43/// Built-in remote task operation.
44#[derive(Debug, Clone, PartialEq, Eq)]
45pub enum RemoteTaskOperation {
46    /// Return the same byte payload.
47    EchoBytes(Vec<u8>),
48    /// Sum `u64` values with wrapping arithmetic.
49    SumU64(Vec<u64>),
50}
51
52/// Owned remote task envelope.
53#[derive(Debug, Clone, PartialEq, Eq)]
54pub struct RemoteTaskEnvelope {
55    /// Task id copied into the result envelope.
56    pub task_id: RemoteTaskId,
57    /// Address where the server sends the result.
58    pub reply_to: RemoteAddress,
59    /// Operation to execute.
60    pub operation: RemoteTaskOperation,
61}
62
63/// Borrowed remote task operation view.
64#[derive(Debug, Clone, Copy, PartialEq, Eq)]
65pub enum RemoteTaskOperationView<'a> {
66    /// Borrowed byte payload.
67    EchoBytes(&'a [u8]),
68    /// Borrowed little-endian u64 list.
69    SumU64(RemoteU64List<'a>),
70}
71
72/// Borrowed `u64` list backed by archive bytes.
73#[derive(Debug, Clone, Copy, PartialEq, Eq)]
74pub struct RemoteU64List<'a> {
75    bytes: &'a [u8],
76    len: usize,
77}
78
79impl RemoteU64List<'_> {
80    /// Number of `u64` values in the list.
81    pub const fn len(self) -> usize {
82        self.len
83    }
84
85    /// Whether the list is empty.
86    pub const fn is_empty(self) -> bool {
87        self.len == 0
88    }
89
90    /// Compute the wrapping sum without materializing an owned vector.
91    pub fn wrapping_sum(self) -> u64 {
92        self.bytes
93            .chunks_exact(core::mem::size_of::<u64>())
94            .fold(0u64, |sum, chunk| {
95                let bytes: [u8; 8] = chunk.try_into().expect("chunk size is fixed");
96                sum.wrapping_add(u64::from_le_bytes(bytes))
97            })
98    }
99}
100
101/// Borrowed remote task envelope view.
102#[derive(Debug, Clone, PartialEq, Eq)]
103pub struct RemoteTaskEnvelopeView<'a> {
104    /// Task id copied into the result envelope.
105    pub task_id: RemoteTaskId,
106    /// Address where the server sends the result.
107    pub reply_to: RemoteAddress,
108    /// Borrowed operation view.
109    pub operation: RemoteTaskOperationView<'a>,
110}
111
112/// Built-in remote task result.
113#[derive(Debug, Clone, PartialEq, Eq)]
114pub enum RemoteTaskOutput {
115    /// Returned byte payload.
116    Bytes(Vec<u8>),
117    /// Returned `u64` value.
118    U64(u64),
119}
120
121/// Owned remote task result envelope.
122#[derive(Debug, Clone, PartialEq, Eq)]
123pub struct RemoteTaskResult {
124    /// Completed task id.
125    pub task_id: RemoteTaskId,
126    /// Operation output.
127    pub output: RemoteTaskOutput,
128}
129
130/// Borrowed remote task result output.
131#[derive(Debug, Clone, Copy, PartialEq, Eq)]
132pub enum RemoteTaskOutputView<'a> {
133    /// Borrowed result bytes.
134    Bytes(&'a [u8]),
135    /// Result integer.
136    U64(u64),
137}
138
139/// Borrowed remote task result envelope view.
140#[derive(Debug, Clone, Copy, PartialEq, Eq)]
141pub struct RemoteTaskResultView<'a> {
142    /// Completed task id.
143    pub task_id: RemoteTaskId,
144    /// Borrowed operation output.
145    pub output: RemoteTaskOutputView<'a>,
146}
147
148/// Remote task client for a fixed server and reply endpoint.
149#[derive(Debug, Clone)]
150pub struct RemoteTaskClient {
151    server: RemoteAddress,
152    reply_to: RemoteAddress,
153}
154
155impl RemoteTaskClient {
156    /// Construct a remote task client.
157    pub fn new(server: RemoteAddress, reply_to: RemoteAddress) -> Self {
158        Self { server, reply_to }
159    }
160
161    /// Send a task, wait for its result, and validate the returned id.
162    pub fn execute(
163        &self,
164        task_id: RemoteTaskId,
165        operation: RemoteTaskOperation,
166    ) -> TransportResult<RemoteTaskResult> {
167        let envelope = RemoteTaskEnvelope {
168            task_id,
169            reply_to: self.reply_to.clone(),
170            operation,
171        };
172
173        let payload = archive_transport_payload::<ServerPayloadRegion, _>(&envelope)?;
174        NetworkTransport {}.send(&Address::Remote(self.server.clone()), payload.into_bytes())?;
175        let bytes = NetworkTransport {}.recv(&Address::Remote(self.reply_to.clone()))?;
176        let payload = TransportPayload::<ServerPayloadRegion>::from_bytes(bytes);
177        let message = ArchivedMessage::<RemoteTaskResult>::from_bytes(payload.into_bytes());
178        let view = message.get()?;
179        if view.task_id != task_id {
180            return Err(TransportError::Closed);
181        }
182
183        Ok(view.into_owned())
184    }
185}
186
187impl ArchiveSerialize for RemoteTaskEnvelope {
188    fn archive_size_hint(&self) -> usize {
189        core::mem::size_of::<u64>()
190            + remote_address_size(&self.reply_to)
191            + operation_size(&self.operation)
192    }
193
194    fn encode_archive(&self, output: &mut Vec<u8>) -> TransportResult<()> {
195        output.extend_from_slice(&self.task_id.get().to_le_bytes());
196        encode_remote_address(&self.reply_to, output)?;
197        match &self.operation {
198            RemoteTaskOperation::EchoBytes(bytes) => {
199                output.push(OP_ECHO_BYTES);
200                encode_len_prefixed_bytes(bytes, output)?;
201            }
202            RemoteTaskOperation::SumU64(values) => {
203                output.push(OP_SUM_U64);
204                let len = u32::try_from(values.len()).map_err(|_| TransportError::Closed)?;
205                output.extend_from_slice(&len.to_le_bytes());
206                for value in values {
207                    output.extend_from_slice(&value.to_le_bytes());
208                }
209            }
210        }
211        Ok(())
212    }
213}
214
215impl ArchiveView for RemoteTaskEnvelope {
216    type Archived<'a> = RemoteTaskEnvelopeView<'a>;
217
218    fn view_archive(bytes: &[u8]) -> TransportResult<Self::Archived<'_>> {
219        let mut cursor = ByteCursor::new(bytes);
220        let task_id = RemoteTaskId::new(u64::from_le_bytes(cursor.read_array()?));
221        let reply_to = cursor.read_remote_address()?;
222        let op = cursor.read_array::<1>()?[0];
223        let operation = match op {
224            OP_ECHO_BYTES => RemoteTaskOperationView::EchoBytes(cursor.read_len_prefixed_bytes()?),
225            OP_SUM_U64 => {
226                let len = usize::try_from(u32::from_le_bytes(cursor.read_array()?))
227                    .map_err(|_| TransportError::Closed)?;
228                let byte_len = len
229                    .checked_mul(core::mem::size_of::<u64>())
230                    .ok_or(TransportError::Closed)?;
231                let payload = cursor.read_exact(byte_len)?;
232                RemoteTaskOperationView::SumU64(RemoteU64List {
233                    bytes: payload,
234                    len,
235                })
236            }
237            _ => return Err(TransportError::Closed),
238        };
239        cursor.finish()?;
240
241        Ok(RemoteTaskEnvelopeView {
242            task_id,
243            reply_to,
244            operation,
245        })
246    }
247}
248
249impl ArchiveSerialize for RemoteTaskResult {
250    fn archive_size_hint(&self) -> usize {
251        core::mem::size_of::<u64>()
252            + 1
253            + match &self.output {
254                RemoteTaskOutput::Bytes(bytes) => core::mem::size_of::<u32>() + bytes.len(),
255                RemoteTaskOutput::U64(_) => core::mem::size_of::<u64>(),
256            }
257    }
258
259    fn encode_archive(&self, output: &mut Vec<u8>) -> TransportResult<()> {
260        output.extend_from_slice(&self.task_id.get().to_le_bytes());
261        match &self.output {
262            RemoteTaskOutput::Bytes(bytes) => {
263                output.push(RESULT_BYTES);
264                encode_len_prefixed_bytes(bytes, output)?;
265            }
266            RemoteTaskOutput::U64(value) => {
267                output.push(RESULT_U64);
268                output.extend_from_slice(&value.to_le_bytes());
269            }
270        }
271        Ok(())
272    }
273}
274
275impl ArchiveView for RemoteTaskResult {
276    type Archived<'a> = RemoteTaskResultView<'a>;
277
278    fn view_archive(bytes: &[u8]) -> TransportResult<Self::Archived<'_>> {
279        let mut cursor = ByteCursor::new(bytes);
280        let task_id = RemoteTaskId::new(u64::from_le_bytes(cursor.read_array()?));
281        let tag = cursor.read_array::<1>()?[0];
282        let output = match tag {
283            RESULT_BYTES => RemoteTaskOutputView::Bytes(cursor.read_len_prefixed_bytes()?),
284            RESULT_U64 => RemoteTaskOutputView::U64(u64::from_le_bytes(cursor.read_array()?)),
285            _ => return Err(TransportError::Closed),
286        };
287        cursor.finish()?;
288
289        Ok(RemoteTaskResultView { task_id, output })
290    }
291}
292
293impl RemoteTaskResultView<'_> {
294    /// Materialize the borrowed result view into an owned result.
295    pub fn into_owned(self) -> RemoteTaskResult {
296        RemoteTaskResult {
297            task_id: self.task_id,
298            output: match self.output {
299                RemoteTaskOutputView::Bytes(bytes) => RemoteTaskOutput::Bytes(bytes.to_vec()),
300                RemoteTaskOutputView::U64(value) => RemoteTaskOutput::U64(value),
301            },
302        }
303    }
304}
305
306pub(super) fn execute_remote_task(request: &RemoteTaskEnvelopeView<'_>) -> RemoteTaskResult {
307    let output = match request.operation {
308        RemoteTaskOperationView::EchoBytes(bytes) => RemoteTaskOutput::Bytes(bytes.to_vec()),
309        RemoteTaskOperationView::SumU64(values) => RemoteTaskOutput::U64(values.wrapping_sum()),
310    };
311
312    RemoteTaskResult {
313        task_id: request.task_id,
314        output,
315    }
316}
317
318fn encode_remote_address(address: &RemoteAddress, output: &mut Vec<u8>) -> TransportResult<()> {
319    encode_len_prefixed_str(&address.host, output)?;
320    output.extend_from_slice(&address.port.to_le_bytes());
321    encode_len_prefixed_str(&address.service, output)
322}
323
324fn remote_address_size(address: &RemoteAddress) -> usize {
325    core::mem::size_of::<u32>()
326        + address.host.len()
327        + core::mem::size_of::<u16>()
328        + core::mem::size_of::<u32>()
329        + address.service.len()
330}
331
332fn operation_size(operation: &RemoteTaskOperation) -> usize {
333    1 + match operation {
334        RemoteTaskOperation::EchoBytes(bytes) => core::mem::size_of::<u32>() + bytes.len(),
335        RemoteTaskOperation::SumU64(values) => {
336            core::mem::size_of::<u32>() + values.len() * core::mem::size_of::<u64>()
337        }
338    }
339}
340
341fn encode_len_prefixed_str(value: &str, output: &mut Vec<u8>) -> TransportResult<()> {
342    encode_len_prefixed_bytes(value.as_bytes(), output)
343}
344
345fn encode_len_prefixed_bytes(value: &[u8], output: &mut Vec<u8>) -> TransportResult<()> {
346    let len = u32::try_from(value.len()).map_err(|_| TransportError::Closed)?;
347    output.extend_from_slice(&len.to_le_bytes());
348    output.extend_from_slice(value);
349    Ok(())
350}
351
352struct ByteCursor<'a> {
353    bytes: &'a [u8],
354    offset: usize,
355}
356
357impl<'a> ByteCursor<'a> {
358    fn new(bytes: &'a [u8]) -> Self {
359        Self { bytes, offset: 0 }
360    }
361
362    /// Read exactly `N` bytes and return them as a fixed-size array.
363    ///
364    /// One generic reader replaces the per-width `read_u16`/`read_u32`/
365    /// `read_u64` clones that previously differed only in the array length and
366    /// the `from_le_bytes` call. `N` is inferred at the call site from the
367    /// integer type it is converted into, so callers stay width-explicit
368    /// (`u32::from_le_bytes(cursor.read_array()?)`) without a per-width
369    /// method.
370    fn read_array<const N: usize>(&mut self) -> TransportResult<[u8; N]> {
371        self.read_exact(N)?
372            .try_into()
373            .map_err(|_| TransportError::Closed)
374    }
375
376    fn read_len_prefixed_bytes(&mut self) -> TransportResult<&'a [u8]> {
377        let len = usize::try_from(u32::from_le_bytes(self.read_array()?))
378            .map_err(|_| TransportError::Closed)?;
379        self.read_exact(len)
380    }
381
382    fn read_len_prefixed_string(&mut self) -> TransportResult<String> {
383        let bytes = self.read_len_prefixed_bytes()?;
384        let value = core::str::from_utf8(bytes).map_err(|_| TransportError::Closed)?;
385        Ok(value.to_string())
386    }
387
388    fn read_remote_address(&mut self) -> TransportResult<RemoteAddress> {
389        let host = self.read_len_prefixed_string()?;
390        let port = u16::from_le_bytes(self.read_array()?);
391        let service = self.read_len_prefixed_string()?;
392        Ok(RemoteAddress {
393            host,
394            port,
395            service,
396        })
397    }
398
399    fn read_exact(&mut self, len: usize) -> TransportResult<&'a [u8]> {
400        let end = self.offset.checked_add(len).ok_or(TransportError::Closed)?;
401        let bytes = self
402            .bytes
403            .get(self.offset..end)
404            .ok_or(TransportError::Closed)?;
405        self.offset = end;
406        Ok(bytes)
407    }
408
409    fn finish(self) -> TransportResult<()> {
410        if self.offset == self.bytes.len() {
411            Ok(())
412        } else {
413            Err(TransportError::Closed)
414        }
415    }
416}
417
418#[cfg(test)]
419mod tests;