Skip to main content

wasi_pg_client/copy/
binary.rs

1//! PostgreSQL binary COPY format encoder.
2//!
3//! The binary COPY format has a specific structure:
4//!
5//! ```text
6//! Header: "PGCOPY\n\xff\r\n\0" (15 bytes)
7//!         + flags (4 bytes, usually 0)
8//!         + header extension length (4 bytes, usually 0)
9//! Tuples: field_count (i16) + [field_length (i32) + field_data (bytes)]*
10//! Trailer: field_count = -1 (i16)
11//! ```
12//!
13//! # Example
14//! ```ignore
15//! use wasi_pg_client::copy::BinaryCopyWriter;
16//!
17//! let mut writer = BinaryCopyWriter::new(2);
18//! let header = writer.header().to_vec();
19//! let row = writer.write_row(&[
20//!     Some(b"42"),
21//!     Some(b"hello"),
22//! ]).to_vec();
23//! let trailer = writer.trailer().to_vec();
24//! ```
25
26/// Writer for PostgreSQL binary COPY format.
27///
28/// This struct helps encode the binary COPY protocol header, individual rows,
29/// and the terminating trailer. Each call to [`write_row`](Self::write_row)
30/// appends to an internal buffer and returns a slice of the newly added bytes.
31#[derive(Debug, Clone)]
32#[non_exhaustive]
33pub struct BinaryCopyWriter {
34    buf: Vec<u8>,
35    column_count: i16,
36    header_written: bool,
37}
38
39impl BinaryCopyWriter {
40    /// Create a new binary COPY writer for the given number of columns.
41    pub fn new(column_count: i16) -> Self {
42        Self {
43            buf: Vec::new(),
44            column_count,
45            header_written: false,
46        }
47    }
48
49    /// Returns the number of columns this writer expects.
50    pub fn column_count(&self) -> i16 {
51        self.column_count
52    }
53
54    /// Generate and return the binary COPY file header.
55    ///
56    /// This should be called once at the start of the COPY stream and the
57    /// returned bytes sent to the server before any row data.
58    ///
59    /// The header consists of:
60    /// - 15-byte magic signature: `PGCOPY\n\xff\r\n\0`
61    /// - 4-byte flags (0)
62    /// - 4-byte header extension length (0)
63    pub fn header(&mut self) -> &[u8] {
64        if self.header_written {
65            return &[];
66        }
67        self.buf.clear();
68        // Magic signature
69        self.buf.extend_from_slice(b"PGCOPY\n\xff\r\n\0");
70        // Flags (0)
71        self.buf.extend_from_slice(&0i32.to_be_bytes());
72        // Header extension length (0)
73        self.buf.extend_from_slice(&0i32.to_be_bytes());
74        self.header_written = true;
75        &self.buf
76    }
77
78    /// Encode a single row and return the bytes.
79    ///
80    /// `values` must have exactly `column_count` elements.
81    /// `None` represents a NULL field.
82    ///
83    /// # Panics
84    /// Panics if `values.len() != column_count`.
85    pub fn write_row(&mut self, values: &[Option<&[u8]>]) -> &[u8] {
86        assert_eq!(
87            values.len() as i16,
88            self.column_count,
89            "value count must match column count"
90        );
91        let start = self.buf.len();
92        self.buf.extend_from_slice(&self.column_count.to_be_bytes());
93        for val in values {
94            match val {
95                Some(data) => {
96                    self.buf
97                        .extend_from_slice(&(data.len() as i32).to_be_bytes());
98                    self.buf.extend_from_slice(data);
99                }
100                None => {
101                    self.buf.extend_from_slice(&(-1i32).to_be_bytes());
102                }
103            }
104        }
105        &self.buf[start..]
106    }
107
108    /// Generate the binary COPY trailer (field_count = -1).
109    ///
110    /// This should be sent after all row data to signal the end of the
111    /// COPY stream.
112    pub fn trailer(&mut self) -> &[u8] {
113        let start = self.buf.len();
114        self.buf.extend_from_slice(&(-1i16).to_be_bytes());
115        &self.buf[start..]
116    }
117
118    /// Consume the writer and return the complete buffer.
119    pub fn into_inner(self) -> Vec<u8> {
120        self.buf
121    }
122
123    /// Return a reference to the accumulated buffer.
124    pub fn buffer(&self) -> &[u8] {
125        &self.buf
126    }
127
128    /// Clear the internal buffer, keeping the allocation.
129    pub fn clear(&mut self) {
130        self.buf.clear();
131        self.header_written = false;
132    }
133}
134
135// ---------------------------------------------------------------------------
136// Tests
137// ---------------------------------------------------------------------------
138
139#[cfg(test)]
140mod tests {
141    use super::*;
142
143    #[test]
144    fn test_binary_header() {
145        let mut writer = BinaryCopyWriter::new(2);
146        let header = writer.header();
147        assert_eq!(header.len(), 11 + 4 + 4); // 11 magic + 4 flags + 4 ext len = 19
148        assert_eq!(&header[..11], b"PGCOPY\n\xff\r\n\0");
149        assert_eq!(&header[11..15], &[0, 0, 0, 0]); // flags = 0
150        assert_eq!(&header[15..19], &[0, 0, 0, 0]); // ext len = 0
151    }
152
153    #[test]
154    fn test_binary_row() {
155        let mut writer = BinaryCopyWriter::new(2);
156        let _header = writer.header();
157        let row = writer.write_row(&[Some(b"42"), Some(b"hello")]);
158
159        // field_count (i16) = 2
160        assert_eq!(&row[..2], &[0, 2]);
161        // field 1 length (i32) = 2, data = "42"
162        assert_eq!(&row[2..6], &[0, 0, 0, 2]);
163        assert_eq!(&row[6..8], b"42");
164        // field 2 length (i32) = 5, data = "hello"
165        assert_eq!(&row[8..12], &[0, 0, 0, 5]);
166        assert_eq!(&row[12..17], b"hello");
167    }
168
169    #[test]
170    fn test_binary_null() {
171        let mut writer = BinaryCopyWriter::new(1);
172        let _header = writer.header();
173        let row = writer.write_row(&[None]);
174
175        // field_count = 1
176        assert_eq!(&row[..2], &[0, 1]);
177        // field length = -1 (NULL)
178        assert_eq!(&row[2..6], &[0xff, 0xff, 0xff, 0xff]);
179    }
180
181    #[test]
182    fn test_binary_trailer() {
183        let mut writer = BinaryCopyWriter::new(1);
184        let trailer = writer.trailer();
185        assert_eq!(trailer.len(), 2);
186        assert_eq!(i16::from_be_bytes([trailer[0], trailer[1]]), -1);
187    }
188
189    #[test]
190    fn test_binary_roundtrip_structure() {
191        let mut writer = BinaryCopyWriter::new(3);
192        let header = writer.header().to_vec();
193        let row1 = writer
194            .write_row(&[Some(b"1"), Some(b"alice"), None])
195            .to_vec();
196        let row2 = writer
197            .write_row(&[Some(b"2"), Some(b"bob"), Some(b"extra")])
198            .to_vec();
199        let trailer = writer.trailer().to_vec();
200
201        let mut all = Vec::new();
202        all.extend_from_slice(&header);
203        all.extend_from_slice(&row1);
204        all.extend_from_slice(&row2);
205        all.extend_from_slice(&trailer);
206
207        // Verify header
208        assert_eq!(&all[..11], b"PGCOPY\n\xff\r\n\0");
209
210        // Verify row 1 starts at offset 19 (after 19-byte header)
211        let row1_offset = 19;
212        assert_eq!(
213            i16::from_be_bytes([all[row1_offset], all[row1_offset + 1]]),
214            3
215        );
216
217        // Verify trailer is last 2 bytes
218        let last_two = &all[all.len() - 2..];
219        assert_eq!(i16::from_be_bytes([last_two[0], last_two[1]]), -1);
220    }
221
222    #[test]
223    #[should_panic(expected = "value count must match column count")]
224    fn test_wrong_column_count_panics() {
225        let mut writer = BinaryCopyWriter::new(2);
226        writer.write_row(&[Some(b"only_one")]);
227    }
228}