Skip to main content

elefant_client/postgres_client/replication/
message_reader.rs

1use super::messages::*;
2use crate::protocol::frame_reader::ByteSliceReader;
3
4pub fn parse_replication_message(data: &[u8]) -> Result<ReplicationMessage<'_>, ReplicationError> {
5    let mut reader = ByteSliceReader::new(data);
6    let msg_type = reader.read_u8()?;
7
8    match msg_type {
9        b'w' => {
10            let start_lsn = Lsn(reader.read_u64()?);
11            let end_lsn = Lsn(reader.read_u64()?);
12            let server_time = reader.read_i64()?;
13            let remaining = reader.read_bytes(data.len() - reader.get_read_bytes())?;
14            Ok(ReplicationMessage::XLogData(XLogData {
15                start_lsn,
16                end_lsn,
17                server_time,
18                data: remaining,
19            }))
20        }
21        b'k' => {
22            let end_lsn = Lsn(reader.read_u64()?);
23            let server_time = reader.read_i64()?;
24            let reply = reader.read_u8()?;
25            Ok(ReplicationMessage::PrimaryKeepalive(PrimaryKeepalive {
26                end_lsn,
27                server_time,
28                reply_requested: reply != 0,
29            }))
30        }
31        _ => Err(ReplicationError::UnknownReplicationMessageType(msg_type)),
32    }
33}
34
35fn parse_tuple_data<'a>(
36    reader: &mut ByteSliceReader<'a>,
37) -> Result<TupleData<'a>, ReplicationError> {
38    let col_count = reader.read_i16()? as usize;
39    let mut columns = Vec::with_capacity(col_count);
40
41    for _ in 0..col_count {
42        let col_type = reader.read_u8()?;
43        match col_type {
44            b'n' => columns.push(TupleColumn::Null),
45            b'u' => columns.push(TupleColumn::Unchanged),
46            b't' => {
47                let len = reader.read_i32()? as usize;
48                let bytes = reader.read_bytes(len)?;
49                let text = String::from_utf8_lossy(bytes);
50                columns.push(TupleColumn::Text(text));
51            }
52            b'b' => {
53                let len = reader.read_i32()? as usize;
54                let bytes = reader.read_bytes(len)?;
55                columns.push(TupleColumn::Binary(bytes));
56            }
57            _ => return Err(ReplicationError::UnknownTupleColumnType(col_type)),
58        }
59    }
60
61    Ok(TupleData { columns })
62}
63
64pub fn parse_pgoutput_message(data: &[u8]) -> Result<PgOutputMessage<'_>, ReplicationError> {
65    let mut reader = ByteSliceReader::new(data);
66    let msg_type = reader.read_u8()?;
67
68    match msg_type {
69        b'B' => {
70            let final_lsn = Lsn(reader.read_u64()?);
71            let commit_timestamp = reader.read_i64()?;
72            let xid = reader.read_i32()? as u32;
73            Ok(PgOutputMessage::Begin(BeginMessage {
74                final_lsn,
75                commit_timestamp,
76                xid,
77            }))
78        }
79        b'C' => {
80            let flags = reader.read_u8()?;
81            let commit_lsn = Lsn(reader.read_u64()?);
82            let end_lsn = Lsn(reader.read_u64()?);
83            let commit_timestamp = reader.read_i64()?;
84            Ok(PgOutputMessage::Commit(CommitMessage {
85                flags,
86                commit_lsn,
87                end_lsn,
88                commit_timestamp,
89            }))
90        }
91        b'R' => {
92            let relation_id = reader.read_i32()? as u32;
93            let namespace = reader.read_null_terminated_string()?;
94            let name = reader.read_null_terminated_string()?;
95            let replica_identity = reader.read_u8()?;
96            let col_count = reader.read_i16()? as usize;
97
98            let mut columns = Vec::with_capacity(col_count);
99            for _ in 0..col_count {
100                let flags = reader.read_u8()?;
101                let col_name = reader.read_null_terminated_string()?;
102                let type_oid = reader.read_i32()? as u32;
103                let type_modifier = reader.read_i32()?;
104                columns.push(RelationColumn {
105                    flags,
106                    name: col_name,
107                    type_oid,
108                    type_modifier,
109                });
110            }
111
112            Ok(PgOutputMessage::Relation(RelationMessage {
113                relation_id,
114                namespace,
115                name,
116                replica_identity,
117                columns,
118            }))
119        }
120        b'I' => {
121            let relation_id = reader.read_i32()? as u32;
122            let _new_marker = reader.read_u8()?; // 'N'
123            let tuple = parse_tuple_data(&mut reader)?;
124            Ok(PgOutputMessage::Insert(InsertMessage {
125                relation_id,
126                tuple,
127            }))
128        }
129        b'U' => {
130            let relation_id = reader.read_i32()? as u32;
131            let marker = reader.read_u8()?;
132
133            let (old_tuple, new_tuple) = match marker {
134                b'K' | b'O' => {
135                    let old = parse_tuple_data(&mut reader)?;
136                    let _new_marker = reader.read_u8()?; // 'N'
137                    let new = parse_tuple_data(&mut reader)?;
138                    (Some(old), new)
139                }
140                b'N' => {
141                    let new = parse_tuple_data(&mut reader)?;
142                    (None, new)
143                }
144                _ => return Err(ReplicationError::UnknownUpdateMarker(marker)),
145            };
146
147            Ok(PgOutputMessage::Update(UpdateMessage {
148                relation_id,
149                old_tuple,
150                new_tuple,
151            }))
152        }
153        b'D' => {
154            let relation_id = reader.read_i32()? as u32;
155            let _marker = reader.read_u8()?; // 'K' or 'O'
156            let old_tuple = parse_tuple_data(&mut reader)?;
157            Ok(PgOutputMessage::Delete(DeleteMessage {
158                relation_id,
159                old_tuple,
160            }))
161        }
162        b'T' => {
163            let num_relations = reader.read_i32()? as usize;
164            let option_bits = reader.read_u8()?;
165            let mut relation_ids = Vec::with_capacity(num_relations);
166            for _ in 0..num_relations {
167                let id = reader.read_i32()? as u32;
168                relation_ids.push(id);
169            }
170            Ok(PgOutputMessage::Truncate(TruncateMessage {
171                option_bits,
172                relation_ids,
173            }))
174        }
175        b'O' => {
176            let origin_lsn = Lsn(reader.read_u64()?);
177            let origin_name = reader.read_null_terminated_string()?;
178            Ok(PgOutputMessage::Origin(OriginMessage {
179                origin_lsn,
180                origin_name,
181            }))
182        }
183        b'Y' => {
184            let type_oid = reader.read_i32()? as u32;
185            let namespace = reader.read_null_terminated_string()?;
186            let name = reader.read_null_terminated_string()?;
187            Ok(PgOutputMessage::Type(TypeMessage {
188                type_oid,
189                namespace,
190                name,
191            }))
192        }
193        b'M' => {
194            let flags = reader.read_u8()?;
195            let transactional = (flags & 1) != 0;
196            let lsn = Lsn(reader.read_u64()?);
197            let prefix = reader.read_null_terminated_string()?;
198            let content_length = reader.read_i32()? as usize;
199            let content = reader.read_bytes(content_length)?;
200            Ok(PgOutputMessage::LogicalDecodingMessage(
201                LogicalDecodingMessage {
202                    transactional,
203                    lsn,
204                    prefix,
205                    content,
206                },
207            ))
208        }
209        b'S' => {
210            let xid = reader.read_i32()? as u32;
211            let first_segment = reader.read_u8()? != 0;
212            Ok(PgOutputMessage::StreamStart(StreamStartMessage {
213                xid,
214                first_segment,
215            }))
216        }
217        b'E' => Ok(PgOutputMessage::StreamStop),
218        b'c' => {
219            let xid = reader.read_i32()? as u32;
220            let flags = reader.read_u8()?;
221            let commit_lsn = Lsn(reader.read_u64()?);
222            let end_lsn = Lsn(reader.read_u64()?);
223            let commit_timestamp = reader.read_i64()?;
224            Ok(PgOutputMessage::StreamCommit(StreamCommitMessage {
225                xid,
226                flags,
227                commit_lsn,
228                end_lsn,
229                commit_timestamp,
230            }))
231        }
232        b'A' => {
233            let xid = reader.read_i32()? as u32;
234            let sub_xid = reader.read_i32()? as u32;
235            Ok(PgOutputMessage::StreamAbort(StreamAbortMessage {
236                xid,
237                sub_xid,
238            }))
239        }
240        _ => {
241            let remaining = data.len() - reader.get_read_bytes();
242            let payload = reader.read_bytes(remaining)?;
243            Ok(PgOutputMessage::Unsupported {
244                msg_type,
245                data: payload,
246            })
247        }
248    }
249}
250
251#[cfg(test)]
252mod tests {
253    use super::*;
254
255    #[test]
256    fn parse_origin_message() {
257        let mut buf = vec![b'O'];
258        buf.extend_from_slice(&42u64.to_be_bytes());
259        buf.extend_from_slice(b"origin_name\0");
260        let msg = parse_pgoutput_message(&buf).unwrap();
261        match msg {
262            PgOutputMessage::Origin(o) => {
263                assert_eq!(o.origin_lsn, Lsn(42));
264                assert_eq!(o.origin_name.as_ref(), "origin_name");
265            }
266            _ => panic!("Expected Origin, got {msg:?}"),
267        }
268    }
269
270    #[test]
271    fn parse_type_message() {
272        let mut buf = vec![b'Y'];
273        buf.extend_from_slice(&100i32.to_be_bytes());
274        buf.extend_from_slice(b"public\0");
275        buf.extend_from_slice(b"my_type\0");
276        let msg = parse_pgoutput_message(&buf).unwrap();
277        match msg {
278            PgOutputMessage::Type(t) => {
279                assert_eq!(t.type_oid, 100);
280                assert_eq!(t.namespace.as_ref(), "public");
281                assert_eq!(t.name.as_ref(), "my_type");
282            }
283            _ => panic!("Expected Type, got {msg:?}"),
284        }
285    }
286
287    #[test]
288    fn parse_logical_decoding_message() {
289        let mut buf = vec![b'M'];
290        buf.push(1); // transactional
291        buf.extend_from_slice(&99u64.to_be_bytes());
292        buf.extend_from_slice(b"my_prefix\0");
293        let content = b"hello world";
294        buf.extend_from_slice(&(content.len() as i32).to_be_bytes());
295        buf.extend_from_slice(content);
296        let msg = parse_pgoutput_message(&buf).unwrap();
297        match msg {
298            PgOutputMessage::LogicalDecodingMessage(m) => {
299                assert!(m.transactional);
300                assert_eq!(m.lsn, Lsn(99));
301                assert_eq!(m.prefix.as_ref(), "my_prefix");
302                assert_eq!(m.content, b"hello world");
303            }
304            _ => panic!("Expected LogicalDecodingMessage, got {msg:?}"),
305        }
306    }
307
308    #[test]
309    fn parse_stream_start_message() {
310        let mut buf = vec![b'S'];
311        buf.extend_from_slice(&42i32.to_be_bytes());
312        buf.push(1); // first segment
313        let msg = parse_pgoutput_message(&buf).unwrap();
314        match msg {
315            PgOutputMessage::StreamStart(s) => {
316                assert_eq!(s.xid, 42);
317                assert!(s.first_segment);
318            }
319            _ => panic!("Expected StreamStart, got {msg:?}"),
320        }
321    }
322
323    #[test]
324    fn parse_stream_stop_message() {
325        let buf = vec![b'E'];
326        let msg = parse_pgoutput_message(&buf).unwrap();
327        assert!(matches!(msg, PgOutputMessage::StreamStop));
328    }
329
330    #[test]
331    fn parse_stream_commit_message() {
332        let mut buf = vec![b'c'];
333        buf.extend_from_slice(&10i32.to_be_bytes()); // xid
334        buf.push(0); // flags
335        buf.extend_from_slice(&100u64.to_be_bytes()); // commit_lsn
336        buf.extend_from_slice(&200u64.to_be_bytes()); // end_lsn
337        buf.extend_from_slice(&999i64.to_be_bytes()); // timestamp
338        let msg = parse_pgoutput_message(&buf).unwrap();
339        match msg {
340            PgOutputMessage::StreamCommit(c) => {
341                assert_eq!(c.xid, 10);
342                assert_eq!(c.flags, 0);
343                assert_eq!(c.commit_lsn, Lsn(100));
344                assert_eq!(c.end_lsn, Lsn(200));
345                assert_eq!(c.commit_timestamp, 999);
346            }
347            _ => panic!("Expected StreamCommit, got {msg:?}"),
348        }
349    }
350
351    #[test]
352    fn parse_stream_abort_message() {
353        let mut buf = vec![b'A'];
354        buf.extend_from_slice(&5i32.to_be_bytes()); // xid
355        buf.extend_from_slice(&6i32.to_be_bytes()); // sub_xid
356        let msg = parse_pgoutput_message(&buf).unwrap();
357        match msg {
358            PgOutputMessage::StreamAbort(a) => {
359                assert_eq!(a.xid, 5);
360                assert_eq!(a.sub_xid, 6);
361            }
362            _ => panic!("Expected StreamAbort, got {msg:?}"),
363        }
364    }
365
366    #[test]
367    fn parse_unsupported_message_type_does_not_error() {
368        // Simulate a two-phase commit 'b' (BeginPrepare) message
369        let data = [b'b', 0, 0, 0, 1, 0, 0, 0, 2];
370        let msg = parse_pgoutput_message(&data).unwrap();
371        match msg {
372            PgOutputMessage::Unsupported { msg_type, data } => {
373                assert_eq!(msg_type, b'b');
374                assert_eq!(data, &[0, 0, 0, 1, 0, 0, 0, 2]);
375            }
376            _ => panic!("Expected Unsupported, got {msg:?}"),
377        }
378    }
379
380    #[test]
381    fn lsn_display_and_parse_roundtrip() {
382        let lsn = Lsn(0x0000_0001_0000_00A0);
383        let s = lsn.to_string();
384        let parsed: Lsn = s.parse().unwrap();
385        assert_eq!(lsn, parsed);
386    }
387}