elefant_client/postgres_client/
copy.rs1use 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 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 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 }
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 }
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}