1use 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 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 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 => {}
43 Err(e) => return Some(Err(ReplicaError::Io(e))),
44 }
45 }
46 }
47
48 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}