diff --git a/tokio-postgres-binary-copy/Cargo.toml b/tokio-postgres-binary-copy/Cargo.toml index e6a70baf..7b3f3b7e 100644 --- a/tokio-postgres-binary-copy/Cargo.toml +++ b/tokio-postgres-binary-copy/Cargo.toml @@ -14,4 +14,3 @@ tokio-postgres = { version = "=0.5.0-alpha.1", default-features = false, path = [dev-dependencies] tokio = "=0.2.0-alpha.6" tokio-postgres = { version = "=0.5.0-alpha.1", path = "../tokio-postgres" } - diff --git a/tokio-postgres-binary-copy/src/lib.rs b/tokio-postgres-binary-copy/src/lib.rs index fa42a84a..99089da3 100644 --- a/tokio-postgres-binary-copy/src/lib.rs +++ b/tokio-postgres-binary-copy/src/lib.rs @@ -1,22 +1,27 @@ -use bytes::{BigEndian, BufMut, ByteOrder, Bytes, BytesMut}; -use futures::{future, Stream}; +use bytes::{BigEndian, BufMut, ByteOrder, Bytes, BytesMut, Buf}; +use futures::{future, ready, Stream}; use parking_lot::Mutex; use pin_project_lite::pin_project; use std::convert::TryFrom; use std::error::Error; use std::future::Future; +use std::ops::Range; use std::pin::Pin; use std::sync::Arc; use std::task::{Context, Poll}; -use tokio_postgres::types::{IsNull, ToSql, Type}; +use tokio_postgres::types::{IsNull, ToSql, Type, FromSql, WrongType}; +use tokio_postgres::CopyStream; +use std::io::Cursor; #[cfg(test)] mod test; const BLOCK_SIZE: usize = 4096; +const MAGIC: &[u8] = b"PGCOPY\n\xff\r\n\0"; +const HEADER_LEN: usize = MAGIC.len() + 4 + 4; pin_project! { - pub struct BinaryCopyStream { + pub struct BinaryCopyInStream { #[pin] future: F, buf: Arc>, @@ -24,27 +29,27 @@ pin_project! { } } -impl BinaryCopyStream +impl BinaryCopyInStream where F: Future>>, { - pub fn new(types: &[Type], write_values: M) -> BinaryCopyStream + pub fn new(types: &[Type], write_values: M) -> BinaryCopyInStream where - M: FnOnce(BinaryCopyWriter) -> F, + M: FnOnce(BinaryCopyInWriter) -> F, { let mut buf = BytesMut::new(); - buf.reserve(11 + 4 + 4); - buf.put_slice(b"PGCOPY\n\xff\r\n\0"); // magic + buf.reserve(HEADER_LEN); + buf.put_slice(MAGIC); // magic buf.put_i32_be(0); // flags buf.put_i32_be(0); // header extension let buf = Arc::new(Mutex::new(buf)); - let writer = BinaryCopyWriter { + let writer = BinaryCopyInWriter { buf: buf.clone(), types: types.to_vec(), }; - BinaryCopyStream { + BinaryCopyInStream { future: write_values(writer), buf, done: false, @@ -52,7 +57,7 @@ where } } -impl Stream for BinaryCopyStream +impl Stream for BinaryCopyInStream where F: Future>>, { @@ -81,12 +86,12 @@ where } // FIXME this should really just take a reference to the buffer, but that requires HKT :( -pub struct BinaryCopyWriter { +pub struct BinaryCopyInWriter { buf: Arc>, types: Vec, } -impl BinaryCopyWriter { +impl BinaryCopyInWriter { pub async fn write( &mut self, values: &[&(dyn ToSql + Send)], @@ -119,7 +124,7 @@ impl BinaryCopyWriter { let mut buf = self.buf.lock(); buf.reserve(2); - buf.put_i16_be(self.types.len() as i16); + buf.put_u16_be(self.types.len() as u16); for (value, type_) in values.zip(&self.types) { let idx = buf.len(); @@ -135,3 +140,131 @@ impl BinaryCopyWriter { Ok(()) } } + +struct Header { + has_oids: bool, +} + +pin_project! { + pub struct BinaryCopyOutStream { + #[pin] + stream: CopyStream, + types: Arc>, + header: Option
, + } +} + +impl BinaryCopyOutStream { + pub fn new(types: &[Type], stream: CopyStream) -> BinaryCopyOutStream { + BinaryCopyOutStream { + stream, + types: Arc::new(types.to_vec()), + header: None, + } + } +} + +impl Stream for BinaryCopyOutStream { + type Item = Result>; + + fn poll_next(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { + let this = self.project(); + + let chunk = match ready!(this.stream.poll_next(cx)) { + Some(Ok(chunk)) => chunk, + Some(Err(e)) => return Poll::Ready(Some(Err(e.into()))), + None => return Poll::Ready(Some(Err("unexpected EOF".into()))), + }; + let mut chunk= Cursor::new(chunk); + + let has_oids = match &this.header { + Some(header) => header.has_oids, + None => { + check_remaining(&chunk, HEADER_LEN)?; + if &chunk.bytes()[..MAGIC.len()] != MAGIC { + return Poll::Ready(Some(Err("invalid magic value".into()))); + } + chunk.advance(MAGIC.len()); + + let flags = chunk.get_i32_be(); + let has_oids = (flags & (1 << 16)) != 0; + + let header_extension = chunk.get_u32_be() as usize; + check_remaining(&chunk, header_extension)?; + chunk.advance(header_extension); + + *this.header = Some(Header { has_oids }); + has_oids + } + }; + + check_remaining(&chunk, 2)?; + let mut len = chunk.get_i16_be(); + if len == -1 { + return Poll::Ready(None); + } + + if has_oids { + len += 1; + } + if len as usize != this.types.len() { + return Poll::Ready(Some(Err("unexpected tuple size".into()))); + } + + let mut ranges = vec![]; + for _ in 0..len { + check_remaining(&chunk, 4)?; + let len = chunk.get_i32_be(); + if len == -1 { + ranges.push(None); + } else { + let len = len as usize; + check_remaining(&chunk, len)?; + let start = chunk.position() as usize; + ranges.push(Some(start..start + len)); + chunk.advance(len); + } + } + + Poll::Ready(Some(Ok(BinaryCopyOutRow { + buf: chunk.into_inner(), + ranges, + types: this.types.clone(), + }))) + } +} + +fn check_remaining(buf: &impl Buf, len: usize) -> Result<(), Box> { + if buf.remaining() < len { + Err("unexpected EOF".into()) + } else { + Ok(()) + } +} + +pub struct BinaryCopyOutRow { + buf: Bytes, + ranges: Vec>>, + types: Arc>, +} + +impl BinaryCopyOutRow { + pub fn try_get<'a, T>(&'a self, idx: usize) -> Result> where T: FromSql<'a> { + let type_ = &self.types[idx]; + if !T::accepts(type_) { + return Err(WrongType::new::(type_.clone()).into()); + } + + match &self.ranges[idx] { + Some(range) => T::from_sql(type_, &self.buf[range.clone()]).map_err(Into::into), + None => T::from_sql_null(type_).map_err(Into::into) + } + } + + pub fn get<'a, T>(&'a self, idx: usize) -> T where T: FromSql<'a> { + match self.try_get(idx) { + Ok(value) => value, + Err(e) => panic!("error retrieving column {}: {}", idx, e), + } + } +} diff --git a/tokio-postgres-binary-copy/src/test.rs b/tokio-postgres-binary-copy/src/test.rs index 486ac581..f4d19351 100644 --- a/tokio-postgres-binary-copy/src/test.rs +++ b/tokio-postgres-binary-copy/src/test.rs @@ -1,6 +1,7 @@ -use crate::BinaryCopyStream; +use crate::{BinaryCopyInStream, BinaryCopyOutStream}; use tokio_postgres::types::Type; use tokio_postgres::{Client, NoTls}; +use futures::TryStreamExt; async fn connect() -> Client { let (client, connection) = @@ -22,7 +23,7 @@ async fn write_basic() { .await .unwrap(); - let stream = BinaryCopyStream::new(&[Type::INT4, Type::TEXT], |mut w| { + let stream = BinaryCopyInStream::new(&[Type::INT4, Type::TEXT], |mut w| { async move { w.write(&[&1i32, &"foobar"]).await?; w.write(&[&2i32, &None::<&str>]).await?; @@ -56,7 +57,7 @@ async fn write_many_rows() { .await .unwrap(); - let stream = BinaryCopyStream::new(&[Type::INT4, Type::TEXT], |mut w| { + let stream = BinaryCopyInStream::new(&[Type::INT4, Type::TEXT], |mut w| { async move { for i in 0..10_000i32 { w.write(&[&i, &format!("the value for {}", i)]).await?; @@ -90,7 +91,7 @@ async fn write_big_rows() { .await .unwrap(); - let stream = BinaryCopyStream::new(&[Type::INT4, Type::BYTEA], |mut w| { + let stream = BinaryCopyInStream::new(&[Type::INT4, Type::BYTEA], |mut w| { async move { for i in 0..2i32 { w.write(&[&i, &vec![i as u8; 128 * 1024]]).await?; @@ -114,3 +115,72 @@ async fn write_big_rows() { assert_eq!(row.get::<_, &[u8]>(1), &*vec![i as u8; 128 * 1024]); } } + +#[tokio::test] +async fn read_basic() { + let client = connect().await; + + client + .batch_execute( + " + CREATE TEMPORARY TABLE foo (id INT, bar TEXT); + INSERT INTO foo (id, bar) VALUES (1, 'foobar'), (2, NULL); + " + ) + .await + .unwrap(); + + let stream = client.copy_out("COPY foo (id, bar) TO STDIN BINARY", &[]).await.unwrap(); + let rows = BinaryCopyOutStream::new(&[Type::INT4, Type::TEXT], stream).try_collect::>().await.unwrap(); + assert_eq!(rows.len(), 2); + + assert_eq!(rows[0].get::(0), 1); + assert_eq!(rows[0].get::>(1), Some("foobar")); + assert_eq!(rows[1].get::(0), 2); + assert_eq!(rows[1].get::>(1), None); +} + +#[tokio::test] +async fn read_many_rows() { + let client = connect().await; + + client + .batch_execute( + " + CREATE TEMPORARY TABLE foo (id INT, bar TEXT); + INSERT INTO foo (id, bar) SELECT i, 'the value for ' || i FROM generate_series(0, 9999) i;" + ) + .await + .unwrap(); + + let stream = client.copy_out("COPY foo (id, bar) TO STDIN BINARY", &[]).await.unwrap(); + let rows = BinaryCopyOutStream::new(&[Type::INT4, Type::TEXT], stream).try_collect::>().await.unwrap(); + assert_eq!(rows.len(), 10_000); + + for (i, row) in rows.iter().enumerate() { + assert_eq!(row.get::(0), i as i32); + assert_eq!(row.get::<&str>(1), format!("the value for {}", i)); + } +} + +#[tokio::test] +async fn read_big_rows() { + let client = connect().await; + + client + .batch_execute("CREATE TEMPORARY TABLE foo (id INT, bar BYTEA)") + .await + .unwrap(); + for i in 0..2i32 { + client.execute("INSERT INTO foo (id, bar) VALUES ($1, $2)", &[&i, &vec![i as u8; 128 * 1024]]).await.unwrap(); + } + + let stream = client.copy_out("COPY foo (id, bar) TO STDIN BINARY", &[]).await.unwrap(); + let rows = BinaryCopyOutStream::new(&[Type::INT4, Type::BYTEA], stream).try_collect::>().await.unwrap(); + assert_eq!(rows.len(), 2); + + for (i, row) in rows.iter().enumerate() { + assert_eq!(row.get::(0), i as i32); + assert_eq!(row.get::<&[u8]>(1), &vec![i as u8; 128 * 1024][..]); + } +}