elefant-client 0.1.0

A pure rust implementation of a postgres client that is independent of the executor runtime
Documentation
use crate::pool::{ConnectionFactory, PoolableClient};
use crate::protocol::{BackendMessage, CopyData};
use crate::{ElefantClientError, PostgresClient, Statement, ToSql};
use tracing::debug;

impl<F: ConnectionFactory> PostgresClient<F> {
    pub async fn copy_out(
        &mut self,
        query: &(impl Statement + ?Sized),
        parameters: &[&dyn ToSql],
    ) -> Result<CopyReader<'_, F>, ElefantClientError> {
        let prepared = query.prepare(self).await?;
        prepared.execute(self, parameters).await?;

        let msg = self.read_next_backend_message().await?;

        match msg {
            BackendMessage::CopyOutResponse(_) => Ok(CopyReader { client: self }),
            _ => Err(ElefantClientError::UnexpectedBackendMessage(format!(
                "Expected CopyOutResponse, got {msg:?}"
            ))),
        }
    }

    pub async fn copy_in(
        &mut self,
        query: &(impl Statement + ?Sized),
        parameters: &[&dyn ToSql],
    ) -> Result<CopyWriter<'_, F>, ElefantClientError> {
        let prepared = query.prepare(self).await?;
        prepared.execute(self, parameters).await?;

        let msg = self.read_next_backend_message().await?;

        match msg {
            BackendMessage::CopyInResponse(_) => Ok(CopyWriter::new(self)),
            _ => Err(ElefantClientError::UnexpectedBackendMessage(format!(
                "Expected CopyInResponse{msg:?}"
            ))),
        }
    }
}

pub struct CopyWriter<'a, F: ConnectionFactory> {
    client: &'a mut PostgresClient<F>,
    data_buffer: Vec<u8>,
    cursor: usize,
}

impl<'a, F: ConnectionFactory> CopyWriter<'a, F> {
    fn new(client: &'a mut PostgresClient<F>) -> Self {
        Self {
            client,
            data_buffer: vec![0; 8192],
            cursor: 0,
        }
    }

    pub async fn write(&mut self, data: &[u8]) -> Result<(), ElefantClientError> {
        if self.data_buffer.len() - self.cursor < data.len() {
            self.write_buffer_content().await?;
        }

        if data.len() > self.data_buffer.len() {
            // Immediately write large messages to avoid having to expand the buffer.
            self.client
                .connection
                .write_frontend_message(&crate::protocol::FrontendMessage::CopyData(CopyData {
                    data,
                }))
                .await?;
        } else {
            self.data_buffer[self.cursor..self.cursor + data.len()].copy_from_slice(data);
            self.cursor += data.len();
        }

        Ok(())
    }

    pub async fn flush(&mut self) -> Result<(), ElefantClientError> {
        self.write_buffer_content().await?;
        self.client.connection.flush().await?;
        Ok(())
    }

    async fn write_buffer_content(&mut self) -> Result<(), ElefantClientError> {
        if self.cursor > 0 {
            self.client
                .connection
                .write_frontend_message(&crate::protocol::FrontendMessage::CopyData(CopyData {
                    data: &self.data_buffer[0..self.cursor],
                }))
                .await?;
            self.cursor = 0;
        }
        Ok(())
    }

    pub async fn end(mut self) -> Result<(), ElefantClientError> {
        self.write_buffer_content().await?;

        self.client
            .connection
            .write_frontend_message(&crate::protocol::FrontendMessage::CopyDone)
            .await?;

        if self.client.sync_required {
            self.client
                .connection
                .write_frontend_message(&crate::protocol::FrontendMessage::Sync)
                .await?;
            self.client.sync_required = false;
        }

        self.client.connection.flush().await?;

        loop {
            let msg = self.client.read_next_backend_message().await?;
            match msg {
                BackendMessage::CommandComplete(_) => {
                    debug!("Copy command completed");
                }
                BackendMessage::ReadyForQuery(_) => {
                    self.client.ready_for_query = true;
                    break;
                }
                _ => {
                    return Err(ElefantClientError::UnexpectedBackendMessage(format!(
                        "Expected CommandComplete or ReadyForQuery, got {msg:?}"
                    )));
                }
            }
        }

        Ok(())
    }
}

pub struct CopyReader<'a, F: ConnectionFactory> {
    client: &'a mut PostgresClient<F>,
}

impl<'a, F: ConnectionFactory> CopyReader<'a, F> {
    pub async fn read(&mut self) -> Result<Option<CopyData<'_>>, ElefantClientError> {
        let msg = self.client.read_next_backend_message().await?;
        match msg {
            BackendMessage::CopyData(cd) => Ok(Some(cd)),
            BackendMessage::CopyDone => Ok(None),
            _ => Err(ElefantClientError::UnexpectedBackendMessage(format!(
                "Expected CopyData or CopyDone, got {msg:?}"
            ))),
        }
    }

    /// Non-consuming cleanup: reads trailing protocol messages after CopyDone.
    /// Call this after `read()` returns `None` to leave the connection in a clean state.
    pub async fn finish(&mut self) -> Result<(), ElefantClientError> {
        loop {
            let msg = self.client.read_next_backend_message().await?;
            match msg {
                BackendMessage::CopyData(_) | BackendMessage::CopyDone => {}
                BackendMessage::CommandComplete(_) | BackendMessage::ReadyForQuery(_) => {
                    self.client.ready_for_query = true;
                    break;
                }
                _ => {
                    return Err(ElefantClientError::UnexpectedBackendMessage(format!(
                        "Expected CommandComplete, ReadyForQuery or CopyData, got {msg:?}"
                    )));
                }
            }
        }

        Ok(())
    }

    pub async fn end(self) -> Result<(), ElefantClientError> {
        loop {
            let msg = self.client.read_next_backend_message().await?;
            match msg {
                BackendMessage::CopyData(_) | BackendMessage::CopyDone => {
                    // Ignore extra copy data messages
                }
                BackendMessage::CommandComplete(_) | BackendMessage::ReadyForQuery(_) => {
                    self.client.ready_for_query = true;
                    break;
                }
                _ => {
                    return Err(ElefantClientError::UnexpectedBackendMessage(format!(
                        "Expected CommandComplete, ReadyForQuery or CopyData, got {msg:?}"
                    )));
                }
            }
        }

        Ok(())
    }

    pub async fn write_to<W: ConnectionFactory>(
        mut self,
        target: &mut CopyWriter<'_, W>,
    ) -> Result<(), ElefantClientError> {
        while let Some(cd) = self.read().await? {
            target.write(cd.data).await?;
        }

        target.flush().await?;
        self.end().await?;

        Ok(())
    }
}

pub struct OwnedCopyReader<F: ConnectionFactory> {
    client: PoolableClient<F>,
}

impl<F: ConnectionFactory> OwnedCopyReader<F> {
    pub async fn new(
        mut client: PoolableClient<F>,
        query: &(impl Statement + ?Sized),
        parameters: &[&dyn ToSql],
    ) -> Result<Self, ElefantClientError> {
        let prepared = query.prepare(&mut *client).await?;
        prepared.execute(&mut *client, parameters).await?;

        {
            let msg = client.read_next_backend_message().await?;
            match msg {
                BackendMessage::CopyOutResponse(_) => {}
                _ => {
                    return Err(ElefantClientError::UnexpectedBackendMessage(format!(
                        "Expected CopyOutResponse, got {msg:?}"
                    )))
                }
            }
        }

        Ok(OwnedCopyReader { client })
    }

    pub async fn read(&mut self) -> Result<Option<CopyData<'_>>, ElefantClientError> {
        let msg = self.client.read_next_backend_message().await?;
        match msg {
            BackendMessage::CopyData(cd) => Ok(Some(cd)),
            BackendMessage::CopyDone => Ok(None),
            _ => Err(ElefantClientError::UnexpectedBackendMessage(format!(
                "Expected CopyData or CopyDone, got {msg:?}"
            ))),
        }
    }

    pub async fn end(mut self) -> Result<(), ElefantClientError> {
        loop {
            let msg = self.client.read_next_backend_message().await?;
            match msg {
                BackendMessage::CopyData(_) | BackendMessage::CopyDone => {
                    // Ignore extra copy data messages
                }
                BackendMessage::CommandComplete(_) | BackendMessage::ReadyForQuery(_) => {
                    self.client.ready_for_query = true;
                    break;
                }
                _ => {
                    return Err(ElefantClientError::UnexpectedBackendMessage(format!(
                        "Expected CommandComplete, ReadyForQuery or CopyData, got {msg:?}"
                    )));
                }
            }
        }

        Ok(())
    }
}

#[cfg(all(test, feature = "tokio"))]
mod tests {
    use crate::test_helpers::get_tokio_test_client;

    #[tokio::test]
    async fn copies_data() {
        let mut source = get_tokio_test_client().await;
        source.execute_non_query_simple(r#"
            drop table if exists source_table;
            create table source_table(id bigint generated by default as identity primary key, value int, txt text);
            insert into source_table(value, txt) values (1, 'one'), (2, 'two'), (3, 'three');
            "#).await.unwrap();

        let mut target = get_tokio_test_client().await;
        target.execute_non_query_simple(r#"
            drop table if exists target_table;
            create table target_table(id bigint generated by default as identity primary key, value int, txt text);
            "#).await.unwrap();

        let copy_out = source
            .copy_out(
                "COPY source_table(id, value, txt) TO STDOUT(format binary)",
                &[],
            )
            .await
            .unwrap();
        let mut copy_in = target
            .copy_in(
                "COPY target_table(id, value, txt) FROM STDIN(format binary)",
                &[],
            )
            .await
            .unwrap();

        copy_out.write_to(&mut copy_in).await.unwrap();
        copy_in.end().await.unwrap();

        let values = target
            .query("select id, value, txt from target_table order by id", &[])
            .await
            .unwrap()
            .collect_to_vec::<(i64, i32, String)>()
            .await
            .unwrap();
        assert_eq!(
            values,
            vec![
                (1, 1, "one".to_string()),
                (2, 2, "two".to_string()),
                (3, 3, "three".to_string())
            ]
        );
    }
}