Skip to main content

elefant_client/postgres_client/
copy.rs

1use crate::pool::{ConnectionFactory, PoolableClient};
2use crate::protocol::{BackendMessage, CopyData};
3use crate::{ElefantClientError, PostgresClient, Statement, ToSql};
4use tracing::debug;
5
6impl<F: ConnectionFactory> PostgresClient<F> {
7    pub async fn copy_out(
8        &mut self,
9        query: &(impl Statement + ?Sized),
10        parameters: &[&dyn ToSql],
11    ) -> Result<CopyReader<'_, F>, ElefantClientError> {
12        let prepared = query.prepare(self).await?;
13        prepared.execute(self, parameters).await?;
14
15        let msg = self.read_next_backend_message().await?;
16
17        match msg {
18            BackendMessage::CopyOutResponse(_) => Ok(CopyReader { client: self }),
19            _ => Err(ElefantClientError::UnexpectedBackendMessage(format!(
20                "Expected CopyOutResponse, got {msg:?}"
21            ))),
22        }
23    }
24
25    pub async fn copy_in(
26        &mut self,
27        query: &(impl Statement + ?Sized),
28        parameters: &[&dyn ToSql],
29    ) -> Result<CopyWriter<'_, F>, ElefantClientError> {
30        let prepared = query.prepare(self).await?;
31        prepared.execute(self, parameters).await?;
32
33        let msg = self.read_next_backend_message().await?;
34
35        match msg {
36            BackendMessage::CopyInResponse(_) => Ok(CopyWriter::new(self)),
37            _ => Err(ElefantClientError::UnexpectedBackendMessage(format!(
38                "Expected CopyInResponse{msg:?}"
39            ))),
40        }
41    }
42}
43
44pub struct CopyWriter<'a, F: ConnectionFactory> {
45    client: &'a mut PostgresClient<F>,
46    data_buffer: Vec<u8>,
47    cursor: usize,
48}
49
50impl<'a, F: ConnectionFactory> CopyWriter<'a, F> {
51    fn new(client: &'a mut PostgresClient<F>) -> Self {
52        Self {
53            client,
54            data_buffer: vec![0; 8192],
55            cursor: 0,
56        }
57    }
58
59    pub async fn write(&mut self, data: &[u8]) -> Result<(), ElefantClientError> {
60        if self.data_buffer.len() - self.cursor < data.len() {
61            self.write_buffer_content().await?;
62        }
63
64        if data.len() > self.data_buffer.len() {
65            // Immediately write large messages to avoid having to expand the buffer.
66            self.client
67                .connection
68                .write_frontend_message(&crate::protocol::FrontendMessage::CopyData(CopyData {
69                    data,
70                }))
71                .await?;
72        } else {
73            self.data_buffer[self.cursor..self.cursor + data.len()].copy_from_slice(data);
74            self.cursor += data.len();
75        }
76
77        Ok(())
78    }
79
80    pub async fn flush(&mut self) -> Result<(), ElefantClientError> {
81        self.write_buffer_content().await?;
82        self.client.connection.flush().await?;
83        Ok(())
84    }
85
86    async fn write_buffer_content(&mut self) -> Result<(), ElefantClientError> {
87        if self.cursor > 0 {
88            self.client
89                .connection
90                .write_frontend_message(&crate::protocol::FrontendMessage::CopyData(CopyData {
91                    data: &self.data_buffer[0..self.cursor],
92                }))
93                .await?;
94            self.cursor = 0;
95        }
96        Ok(())
97    }
98
99    pub async fn end(mut self) -> Result<(), ElefantClientError> {
100        self.write_buffer_content().await?;
101
102        self.client
103            .connection
104            .write_frontend_message(&crate::protocol::FrontendMessage::CopyDone)
105            .await?;
106
107        if self.client.sync_required {
108            self.client
109                .connection
110                .write_frontend_message(&crate::protocol::FrontendMessage::Sync)
111                .await?;
112            self.client.sync_required = false;
113        }
114
115        self.client.connection.flush().await?;
116
117        loop {
118            let msg = self.client.read_next_backend_message().await?;
119            match msg {
120                BackendMessage::CommandComplete(_) => {
121                    debug!("Copy command completed");
122                }
123                BackendMessage::ReadyForQuery(_) => {
124                    self.client.ready_for_query = true;
125                    break;
126                }
127                _ => {
128                    return Err(ElefantClientError::UnexpectedBackendMessage(format!(
129                        "Expected CommandComplete or ReadyForQuery, got {msg:?}"
130                    )));
131                }
132            }
133        }
134
135        Ok(())
136    }
137}
138
139pub struct CopyReader<'a, F: ConnectionFactory> {
140    client: &'a mut PostgresClient<F>,
141}
142
143impl<'a, F: ConnectionFactory> CopyReader<'a, F> {
144    pub async fn read(&mut self) -> Result<Option<CopyData<'_>>, ElefantClientError> {
145        let msg = self.client.read_next_backend_message().await?;
146        match msg {
147            BackendMessage::CopyData(cd) => Ok(Some(cd)),
148            BackendMessage::CopyDone => Ok(None),
149            _ => Err(ElefantClientError::UnexpectedBackendMessage(format!(
150                "Expected CopyData or CopyDone, got {msg:?}"
151            ))),
152        }
153    }
154
155    /// Non-consuming cleanup: reads trailing protocol messages after CopyDone.
156    /// Call this after `read()` returns `None` to leave the connection in a clean state.
157    pub async fn finish(&mut self) -> Result<(), ElefantClientError> {
158        loop {
159            let msg = self.client.read_next_backend_message().await?;
160            match msg {
161                BackendMessage::CopyData(_) | BackendMessage::CopyDone => {}
162                BackendMessage::CommandComplete(_) | BackendMessage::ReadyForQuery(_) => {
163                    self.client.ready_for_query = true;
164                    break;
165                }
166                _ => {
167                    return Err(ElefantClientError::UnexpectedBackendMessage(format!(
168                        "Expected CommandComplete, ReadyForQuery or CopyData, got {msg:?}"
169                    )));
170                }
171            }
172        }
173
174        Ok(())
175    }
176
177    pub async fn end(self) -> Result<(), ElefantClientError> {
178        loop {
179            let msg = self.client.read_next_backend_message().await?;
180            match msg {
181                BackendMessage::CopyData(_) | BackendMessage::CopyDone => {
182                    // Ignore extra copy data messages
183                }
184                BackendMessage::CommandComplete(_) | BackendMessage::ReadyForQuery(_) => {
185                    self.client.ready_for_query = true;
186                    break;
187                }
188                _ => {
189                    return Err(ElefantClientError::UnexpectedBackendMessage(format!(
190                        "Expected CommandComplete, ReadyForQuery or CopyData, got {msg:?}"
191                    )));
192                }
193            }
194        }
195
196        Ok(())
197    }
198
199    pub async fn write_to<W: ConnectionFactory>(
200        mut self,
201        target: &mut CopyWriter<'_, W>,
202    ) -> Result<(), ElefantClientError> {
203        while let Some(cd) = self.read().await? {
204            target.write(cd.data).await?;
205        }
206
207        target.flush().await?;
208        self.end().await?;
209
210        Ok(())
211    }
212}
213
214pub struct OwnedCopyReader<F: ConnectionFactory> {
215    client: PoolableClient<F>,
216}
217
218impl<F: ConnectionFactory> OwnedCopyReader<F> {
219    pub async fn new(
220        mut client: PoolableClient<F>,
221        query: &(impl Statement + ?Sized),
222        parameters: &[&dyn ToSql],
223    ) -> Result<Self, ElefantClientError> {
224        let prepared = query.prepare(&mut *client).await?;
225        prepared.execute(&mut *client, parameters).await?;
226
227        {
228            let msg = client.read_next_backend_message().await?;
229            match msg {
230                BackendMessage::CopyOutResponse(_) => {}
231                _ => {
232                    return Err(ElefantClientError::UnexpectedBackendMessage(format!(
233                        "Expected CopyOutResponse, got {msg:?}"
234                    )))
235                }
236            }
237        }
238
239        Ok(OwnedCopyReader { client })
240    }
241
242    pub async fn read(&mut self) -> Result<Option<CopyData<'_>>, ElefantClientError> {
243        let msg = self.client.read_next_backend_message().await?;
244        match msg {
245            BackendMessage::CopyData(cd) => Ok(Some(cd)),
246            BackendMessage::CopyDone => Ok(None),
247            _ => Err(ElefantClientError::UnexpectedBackendMessage(format!(
248                "Expected CopyData or CopyDone, got {msg:?}"
249            ))),
250        }
251    }
252
253    pub async fn end(mut self) -> Result<(), ElefantClientError> {
254        loop {
255            let msg = self.client.read_next_backend_message().await?;
256            match msg {
257                BackendMessage::CopyData(_) | BackendMessage::CopyDone => {
258                    // Ignore extra copy data messages
259                }
260                BackendMessage::CommandComplete(_) | BackendMessage::ReadyForQuery(_) => {
261                    self.client.ready_for_query = true;
262                    break;
263                }
264                _ => {
265                    return Err(ElefantClientError::UnexpectedBackendMessage(format!(
266                        "Expected CommandComplete, ReadyForQuery or CopyData, got {msg:?}"
267                    )));
268                }
269            }
270        }
271
272        Ok(())
273    }
274}
275
276#[cfg(all(test, feature = "tokio"))]
277mod tests {
278    use crate::test_helpers::get_tokio_test_client;
279
280    #[tokio::test]
281    async fn copies_data() {
282        let mut source = get_tokio_test_client().await;
283        source.execute_non_query_simple(r#"
284            drop table if exists source_table;
285            create table source_table(id bigint generated by default as identity primary key, value int, txt text);
286            insert into source_table(value, txt) values (1, 'one'), (2, 'two'), (3, 'three');
287            "#).await.unwrap();
288
289        let mut target = get_tokio_test_client().await;
290        target.execute_non_query_simple(r#"
291            drop table if exists target_table;
292            create table target_table(id bigint generated by default as identity primary key, value int, txt text);
293            "#).await.unwrap();
294
295        let copy_out = source
296            .copy_out(
297                "COPY source_table(id, value, txt) TO STDOUT(format binary)",
298                &[],
299            )
300            .await
301            .unwrap();
302        let mut copy_in = target
303            .copy_in(
304                "COPY target_table(id, value, txt) FROM STDIN(format binary)",
305                &[],
306            )
307            .await
308            .unwrap();
309
310        copy_out.write_to(&mut copy_in).await.unwrap();
311        copy_in.end().await.unwrap();
312
313        let values = target
314            .query("select id, value, txt from target_table order by id", &[])
315            .await
316            .unwrap()
317            .collect_to_vec::<(i64, i32, String)>()
318            .await
319            .unwrap();
320        assert_eq!(
321            values,
322            vec![
323                (1, 1, "one".to_string()),
324                (2, 2, "two".to_string()),
325                (3, 3, "three".to_string())
326            ]
327        );
328    }
329}