diff --git a/src/stmt.rs b/src/stmt.rs index 9ebbb19d..0e041111 100644 --- a/src/stmt.rs +++ b/src/stmt.rs @@ -4,7 +4,7 @@ use debug_builders::DebugStruct; use std::cell::{Cell, RefMut}; use std::collections::VecDeque; use std::fmt; -use std::io::{self, Cursor, BufRead, Read}; +use std::io::{self, Read, Write}; use error::{Error, DbError}; use types::{SessionInfo, Type, ToSql, IsNull}; @@ -379,8 +379,8 @@ impl<'conn> Statement<'conn> { Ok(num) } - /// Executes a `COPY TO STDOUT` statement, returning a `Read`er of the - /// resulting data. + /// Executes a `COPY TO STDOUT` statement, passing the resulting data to + /// the provided writer and returning the number of rows received. /// /// See the [Postgres documentation](http://www.postgresql.org/docs/9.4/static/sql-copy.html) /// for details on the data format. @@ -388,28 +388,20 @@ impl<'conn> Statement<'conn> { /// If the statement is not a `COPY TO STDOUT` statement it will still be /// executed and this method will return an error. /// - /// # Warning - /// - /// The underlying connection may not be used while the returned `Read`er - /// exists. Any attempt to do so will panic. - /// /// # Examples /// /// ```rust,no_run - /// # use std::io::Read; /// # use postgres::{Connection, SslMode}; /// # let conn = Connection::connect("", &SslMode::None).unwrap(); /// conn.batch_execute(" /// CREATE TABLE people (id INT PRIMARY KEY, name VARCHAR); /// INSERT INTO people (id, name) VALUES (1, 'john'), (2, 'jane');").unwrap(); /// let stmt = conn.prepare("COPY people TO STDOUT").unwrap(); - /// let mut r = stmt.copy_out(&[]).unwrap(); /// let mut buf = vec![]; - /// r.read_to_end(&mut buf).unwrap(); - /// r.finish().unwrap(); + /// let mut r = stmt.copy_out(&[], &mut buf).unwrap(); /// assert_eq!(buf, b"1\tjohn\n2\tjane\n"); /// ``` - pub fn copy_out<'a>(&'a self, params: &[&ToSql]) -> Result> { + pub fn copy_out<'a, W: WriteWithInfo>(&'a self, params: &[&ToSql], w: &mut W) -> Result { try!(self.inner_execute("", 0, params)); let mut conn = self.conn.conn.borrow_mut(); @@ -448,15 +440,57 @@ impl<'conn> Statement<'conn> { } }; - Ok(CopyOutReader { - info: CopyInfo { - conn: conn, - format: Format::from_u16(format as u16), - column_formats: column_formats.iter().map(|&f| Format::from_u16(f)).collect(), - }, - buf: Cursor::new(vec![]), - finished: false, - }) + let mut info = CopyInfo { + conn: conn, + format: Format::from_u16(format as u16), + column_formats: column_formats.iter().map(|&f| Format::from_u16(f)).collect(), + }; + + let count; + loop { + match try!(info.conn.read_message()) { + BCopyData { data } => { + let mut data = &data[..]; + while !data.is_empty() { + match w.write_with_info(data, &info) { + Ok(n) => data = &data[n..], + Err(e) => { + loop { + match try!(info.conn.read_message()) { + ReadyForQuery { .. } => return Err(Error::IoError(e)), + _ => {} + } + } + } + } + } + } + BCopyDone => {}, + CommandComplete { tag } => { + count = util::parse_update_count(tag); + break; + } + ErrorResponse { fields } => { + loop { + match try!(info.conn.read_message()) { + ReadyForQuery { .. } => return DbError::new(fields), + _ => {} + } + } + } + _ => { + loop { + match try!(info.conn.read_message()) { + ReadyForQuery { .. } => return Err(Error::IoError(bad_response())), + _ => {} + } + } + } + } + } + + try!(info.conn.wait_for_ready()); + Ok(count) } /// Consumes the statement, clearing it from the Postgres session. @@ -539,6 +573,20 @@ impl ReadWithInfo for R { } } +/// Like `Write` except that a `CopyInfo` object is provided as well. +/// +/// All types that implement `Write` also implement this trait. +pub trait WriteWithInfo { + /// Like `Write::write`. + fn write_with_info(&mut self, buf: &[u8], info: &CopyInfo) -> io::Result; +} + +impl WriteWithInfo for W { + fn write_with_info(&mut self, buf: &[u8], _: &CopyInfo) -> io::Result { + self.write(buf) + } +} + impl Column { /// The name of the column. pub fn name(&self) -> &str { @@ -568,95 +616,3 @@ impl Format { } } } - -/// A `Read`er for data from `COPY TO STDOUT` queries. -/// -/// # Warning -/// -/// The underlying connection may not be used while a `CopyOutReader` exists. -/// Any attempt to do so will panic. -pub struct CopyOutReader<'a> { - info: CopyInfo<'a>, - buf: Cursor>, - finished: bool, -} - -impl<'a> Drop for CopyOutReader<'a> { - fn drop(&mut self) { - let _ = self.finish_inner(); - } -} - -impl<'a> CopyOutReader<'a> { - /// Returns the `CopyInfo` for the current operation. - pub fn info(&self) -> &CopyInfo { - &self.info - } - - /// Consumes the `CopyOutReader`, throwing away any unread data. - /// - /// Functionally equivalent to `CopyOutReader`'s `Drop` implementation, - /// except that it returns any error encountered to the caller. - pub fn finish(mut self) -> Result<()> { - self.finish_inner() - } - - fn finish_inner(&mut self) -> Result<()> { - while !self.finished { - let pos = self.buf.get_ref().len() as u64; - self.buf.set_position(pos); - try!(self.ensure_filled()); - } - Ok(()) - } - - fn ensure_filled(&mut self) -> Result<()> { - if self.finished || self.buf.position() != self.buf.get_ref().len() as u64 { - return Ok(()); - } - - match try!(self.info.conn.read_message()) { - BCopyData { data } => self.buf = Cursor::new(data), - BCopyDone => { - self.finished = true; - match try!(self.info.conn.read_message()) { - CommandComplete { .. } => {} - _ => { - self.info.conn.desynchronized = true; - return Err(Error::IoError(bad_response())); - } - } - try!(self.info.conn.wait_for_ready()); - } - ErrorResponse { fields } => { - self.finished = true; - try!(self.info.conn.wait_for_ready()); - return DbError::new(fields); - } - _ => { - self.info.conn.desynchronized = true; - return Err(Error::IoError(bad_response())); - } - } - - Ok(()) - } -} - -impl<'a> Read for CopyOutReader<'a> { - fn read(&mut self, buf: &mut [u8]) -> io::Result { - try!(self.ensure_filled()); - self.buf.read(buf) - } -} - -impl<'a> BufRead for CopyOutReader<'a> { - fn fill_buf(&mut self) -> io::Result<&[u8]> { - try!(self.ensure_filled()); - self.buf.fill_buf() - } - - fn consume(&mut self, amt: usize) { - self.buf.consume(amt) - } -} diff --git a/tests/test.rs b/tests/test.rs index 738bbf9c..c3a83ce0 100644 --- a/tests/test.rs +++ b/tests/test.rs @@ -778,26 +778,10 @@ fn test_copy_out() { CREATE TEMPORARY TABLE foo (id INT); INSERT INTO foo (id) VALUES (0), (1), (2), (3)")); let stmt = or_panic!(conn.prepare("COPY (SELECT id FROM foo ORDER BY id) TO STDOUT")); - let mut reader = or_panic!(stmt.copy_out(&[])); - let mut out = vec![]; - or_panic!(reader.read_to_end(&mut out)); - assert_eq!(out, b"0\n1\n2\n3\n"); - drop(reader); - or_panic!(conn.batch_execute("SELECT 1")); -} - -#[test] -fn test_copy_out_partial_read() { - let conn = or_panic!(Connection::connect("postgres://postgres@localhost", &SslMode::None)); - or_panic!(conn.batch_execute(" - CREATE TEMPORARY TABLE foo (id INT); - INSERT INTO foo (id) VALUES (0), (1), (2), (3)")); - let stmt = or_panic!(conn.prepare("COPY (SELECT id FROM foo ORDER BY id) TO STDOUT")); - let mut reader = or_panic!(stmt.copy_out(&[])); - let mut out = vec![]; - or_panic!(reader.by_ref().take(5).read_to_end(&mut out)); - assert_eq!(out, b"0\n1\n2"); - drop(reader); + let mut buf = vec![]; + let count = or_panic!(stmt.copy_out(&[], &mut buf)); + assert_eq!(count, 4); + assert_eq!(buf, b"0\n1\n2\n3\n"); or_panic!(conn.batch_execute("SELECT 1")); }