Skip to main content

kevy_replicate/
replica_decode.rs

1//! Snapshot-aware event decoding for [`crate::replica::ReplicaClient`]
2//! — split from `replica.rs` to keep that file under the 500-LOC
3//! project ceiling. The state-machine helpers live here; the type
4//! definitions, `connect`, and `next_frame` stay in `replica.rs`.
5
6use crate::replica::{DecodedFrame, ReplicaClient, ReplicaError, ReplicaEvent};
7use crate::wire::{
8    SnapshotMarker, WireError, decode_frame, decode_snapshot_chunk, decode_snapshot_marker,
9};
10use std::io::{self, Read};
11
12impl ReplicaClient {
13    /// Snapshot-aware iterator step. Returns one [`ReplicaEvent`] per
14    /// call — a live `Frame`, a `SnapshotBegin`/`SnapshotChunk`/
15    /// `SnapshotEnd`, or one of the [`ReplicaError`] variants.
16    /// Returns `None` on clean peer EOF.
17    ///
18    /// Snapshot bookkeeping:
19    /// - Entering `SnapshotBegin` sets `in_snapshot = true`; chunk
20    ///   bytes are valid until `SnapshotEnd`.
21    /// - `SnapshotEnd { ack_offset, routed: false }` sets `expected_offset =
22    ///   ack_offset` (so the next live `Frame` has no gap) and
23    ///   clears `in_snapshot`.
24    /// - Live `*2\r\n` bytes during a snapshot return
25    ///   [`ReplicaError::UnexpectedInSnapshot`] (v1.18 forbids
26    ///   interleaving — see `docs/snapshot.md`).
27    pub fn next_event(&mut self) -> Option<Result<ReplicaEvent, ReplicaError>> {
28        loop {
29            if let Some(result) = self.try_decode_one_event() {
30                return Some(result);
31            }
32            // Need more bytes off the socket.
33            let mut chunk = [0u8; 4096];
34            match self.sock.read(&mut chunk) {
35                Ok(0) => {
36                    if self.cursor < self.buf.len() {
37                        return Some(Err(ReplicaError::Truncated));
38                    }
39                    return None;
40                }
41                Ok(n) => self.buf.extend_from_slice(&chunk[..n]),
42                Err(e) if e.kind() == io::ErrorKind::Interrupted => continue,
43                Err(e) => return Some(Err(ReplicaError::Io(e))),
44            }
45        }
46    }
47
48    /// Try to decode one event from the buffered bytes. Returns
49    /// `None` when more bytes are needed (the loop in [`Self::next_event`]
50    /// will read more). Split out so the outer loop stays tiny + the
51    /// per-event dispatch fits the project's 50-LOC-fn rule.
52    fn try_decode_one_event(&mut self) -> Option<Result<ReplicaEvent, ReplicaError>> {
53        if self.cursor >= self.buf.len() {
54            return None;
55        }
56        let first = self.buf[self.cursor];
57        match first {
58            b'+' => self.try_decode_snapshot_marker(),
59            b'$' if self.in_snapshot => self.try_decode_snapshot_chunk(),
60            b'*' if self.in_snapshot => {
61                Some(Err(ReplicaError::UnexpectedInSnapshot))
62            }
63            b'*' => self.try_decode_live_frame(),
64            _ => Some(Err(ReplicaError::Frame(WireError::BadEnvelope))),
65        }
66    }
67
68    fn try_decode_live_frame(&mut self) -> Option<Result<ReplicaEvent, ReplicaError>> {
69        match decode_frame(&self.buf[self.cursor..]) {
70            Ok((offset, argv, used)) => {
71                self.cursor += used;
72                self.maybe_compact_buf();
73                if offset != self.expected_offset {
74                    return Some(Err(ReplicaError::OffsetGap {
75                        expected: self.expected_offset,
76                        got: offset,
77                    }));
78                }
79                self.expected_offset = self.expected_offset.saturating_add(1);
80                Some(Ok(ReplicaEvent::Frame(DecodedFrame { offset, argv })))
81            }
82            Err(WireError::Truncated) => None,
83            Err(e) => Some(Err(ReplicaError::Frame(e))),
84        }
85    }
86
87    fn try_decode_snapshot_marker(&mut self) -> Option<Result<ReplicaEvent, ReplicaError>> {
88        match decode_snapshot_marker(&self.buf[self.cursor..]) {
89            Ok(Some((SnapshotMarker::Begin, used))) => {
90                self.cursor += used;
91                self.maybe_compact_buf();
92                self.in_snapshot = true;
93                Some(Ok(ReplicaEvent::SnapshotBegin))
94            }
95            Ok(Some((SnapshotMarker::Ping { generation, next_offset }, used))) => {
96                self.cursor += used;
97                self.maybe_compact_buf();
98                Some(Ok(ReplicaEvent::Ping { generation, primary_offset: next_offset }))
99            }
100            Ok(Some((SnapshotMarker::End(ack_offset), used))) => {
101                self.cursor += used;
102                self.maybe_compact_buf();
103                self.in_snapshot = false;
104                self.expected_offset = ack_offset;
105                Some(Ok(ReplicaEvent::SnapshotEnd { ack_offset }))
106            }
107            Ok(None) => Some(Err(ReplicaError::Frame(WireError::BadEnvelope))),
108            Err(WireError::Truncated) => None,
109            Err(e) => Some(Err(ReplicaError::Frame(e))),
110        }
111    }
112
113    fn try_decode_snapshot_chunk(&mut self) -> Option<Result<ReplicaEvent, ReplicaError>> {
114        match decode_snapshot_chunk(&self.buf[self.cursor..]) {
115            Ok((chunk, used)) => {
116                let owned = chunk.to_vec();
117                self.cursor += used;
118                self.maybe_compact_buf();
119                Some(Ok(ReplicaEvent::SnapshotChunk(owned)))
120            }
121            Err(WireError::Truncated) => None,
122            Err(e) => Some(Err(ReplicaError::Frame(e))),
123        }
124    }
125}
126
127#[cfg(test)]
128mod tests {
129    use crate::replica::{ReplicaClient, ReplicaError, ReplicaEvent};
130    use crate::wire::{encode_frame, encode_snapshot_begin, encode_snapshot_chunk, encode_snapshot_end};
131    use kevy_resp::Argv;
132    use std::io::Write;
133    use std::net::{TcpListener, TcpStream};
134    use std::thread;
135
136    fn tcp_pair() -> (TcpStream, TcpStream) {
137        let listener = TcpListener::bind("127.0.0.1:0").unwrap();
138        let addr = listener.local_addr().unwrap();
139        let client = TcpStream::connect(addr).unwrap();
140        let (server, _) = listener.accept().unwrap();
141        (server, client)
142    }
143
144    fn argv_for(args: &[&[u8]]) -> Argv {
145        let mut a = Argv::default();
146        for arg in args {
147            a.push(arg);
148        }
149        a
150    }
151
152    #[test]
153    fn next_event_snapshot_path_begin_chunks_end_then_frame() {
154        let (mut srv, cli) = tcp_pair();
155        thread::spawn(move || {
156            srv.write_all(&encode_snapshot_begin()).unwrap();
157            srv.write_all(&encode_snapshot_chunk(b"hello-snapshot")).unwrap();
158            srv.write_all(&encode_snapshot_chunk(b"more-snapshot-bytes")).unwrap();
159            srv.write_all(&encode_snapshot_end(42)).unwrap();
160            srv.write_all(&encode_frame(42, &argv_for(&[b"SET", b"k", b"v"]))).unwrap();
161            std::thread::sleep(std::time::Duration::from_millis(50));
162            drop(srv);
163        });
164        let mut client = ReplicaClient::from_socket_for_test(cli, 0);
165
166        assert!(matches!(client.next_event(), Some(Ok(ReplicaEvent::SnapshotBegin))));
167        match client.next_event() {
168            Some(Ok(ReplicaEvent::SnapshotChunk(bytes))) => {
169                assert_eq!(bytes, b"hello-snapshot");
170            }
171            other => panic!("expected SnapshotChunk, got {other:?}"),
172        }
173        match client.next_event() {
174            Some(Ok(ReplicaEvent::SnapshotChunk(bytes))) => {
175                assert_eq!(bytes, b"more-snapshot-bytes");
176            }
177            other => panic!("expected SnapshotChunk, got {other:?}"),
178        }
179        match client.next_event() {
180            Some(Ok(ReplicaEvent::SnapshotEnd { ack_offset })) => assert_eq!(ack_offset, 42),
181            other => panic!("expected SnapshotEnd, got {other:?}"),
182        }
183        assert_eq!(client.expected_offset(), 42);
184        match client.next_event() {
185            Some(Ok(ReplicaEvent::Frame(f))) => {
186                assert_eq!(f.offset, 42);
187                assert_eq!(f.argv, argv_for(&[b"SET", b"k", b"v"]));
188            }
189            other => panic!("expected Frame, got {other:?}"),
190        }
191    }
192
193    #[test]
194    fn next_event_live_frame_during_snapshot_is_unexpected() {
195        let (mut srv, cli) = tcp_pair();
196        thread::spawn(move || {
197            srv.write_all(&encode_snapshot_begin()).unwrap();
198            srv.write_all(&encode_snapshot_chunk(b"first")).unwrap();
199            srv.write_all(&encode_frame(0, &argv_for(&[b"PING"]))).unwrap();
200            std::thread::sleep(std::time::Duration::from_millis(50));
201            drop(srv);
202        });
203        let mut client = ReplicaClient::from_socket_for_test(cli, 0);
204        assert!(matches!(client.next_event(), Some(Ok(ReplicaEvent::SnapshotBegin))));
205        assert!(matches!(client.next_event(), Some(Ok(ReplicaEvent::SnapshotChunk(_)))));
206        assert!(matches!(
207            client.next_event(),
208            Some(Err(ReplicaError::UnexpectedInSnapshot))
209        ));
210    }
211
212    #[test]
213    fn next_frame_returns_snapshot_in_progress_when_snapshot_starts() {
214        let (mut srv, cli) = tcp_pair();
215        thread::spawn(move || {
216            srv.write_all(&encode_snapshot_begin()).unwrap();
217            std::thread::sleep(std::time::Duration::from_millis(50));
218            drop(srv);
219        });
220        let mut client = ReplicaClient::from_socket_for_test(cli, 0);
221        assert!(matches!(
222            client.next_frame(),
223            Some(Err(ReplicaError::SnapshotInProgress))
224        ));
225    }
226
227    #[test]
228    fn live_frame_path_via_next_event() {
229        let (mut srv, cli) = tcp_pair();
230        thread::spawn(move || {
231            srv.write_all(&encode_frame(0, &argv_for(&[b"SET", b"a", b"1"]))).unwrap();
232            srv.write_all(&encode_frame(1, &argv_for(&[b"SET", b"b", b"2"]))).unwrap();
233            std::thread::sleep(std::time::Duration::from_millis(50));
234            drop(srv);
235        });
236        let mut client = ReplicaClient::from_socket_for_test(cli, 0);
237        for expected_off in 0..2 {
238            match client.next_event() {
239                Some(Ok(ReplicaEvent::Frame(f))) => assert_eq!(f.offset, expected_off),
240                other => panic!("expected Frame {expected_off}, got {other:?}"),
241            }
242        }
243        assert_eq!(client.expected_offset(), 2);
244    }
245
246    #[test]
247    fn snapshot_end_with_zero_offset_handled() {
248        let (mut srv, cli) = tcp_pair();
249        thread::spawn(move || {
250            srv.write_all(&encode_snapshot_begin()).unwrap();
251            srv.write_all(&encode_snapshot_end(0)).unwrap();
252            std::thread::sleep(std::time::Duration::from_millis(50));
253            drop(srv);
254        });
255        let mut client = ReplicaClient::from_socket_for_test(cli, 0);
256        assert!(matches!(client.next_event(), Some(Ok(ReplicaEvent::SnapshotBegin))));
257        match client.next_event() {
258            Some(Ok(ReplicaEvent::SnapshotEnd { ack_offset })) => assert_eq!(ack_offset, 0),
259            other => panic!("expected SnapshotEnd, got {other:?}"),
260        }
261        assert_eq!(client.expected_offset(), 0);
262    }
263}