diff --git a/postgres-protocol/src/authentication/sasl.rs b/postgres-protocol/src/authentication/sasl.rs index 617a288f..dfbb70b2 100644 --- a/postgres-protocol/src/authentication/sasl.rs +++ b/postgres-protocol/src/authentication/sasl.rs @@ -141,7 +141,7 @@ pub struct ScramSha256 { impl ScramSha256 { /// Constructs a new instance which will use the provided password for authentication. - pub fn new(password: &[u8], channel_binding: ChannelBinding) -> io::Result { + pub fn new(password: &[u8], channel_binding: ChannelBinding) -> ScramSha256 { // rand 0.5's ThreadRng is cryptographically secure let mut rng = rand::thread_rng(); let nonce = (0..NONCE_LENGTH) @@ -151,25 +151,20 @@ impl ScramSha256 { v = 0x7e } v as char - }) - .collect::(); + }).collect::(); ScramSha256::new_inner(password, channel_binding, nonce) } - fn new_inner( - password: &[u8], - channel_binding: ChannelBinding, - nonce: String, - ) -> io::Result { - Ok(ScramSha256 { + fn new_inner(password: &[u8], channel_binding: ChannelBinding, nonce: String) -> ScramSha256 { + ScramSha256 { message: format!("{}n=,r={}", channel_binding.gs2_header(), nonce), state: State::Update { nonce, password: normalize(password), channel_binding: channel_binding, }, - }) + } } /// Returns the message which should be sent to the backend in an `SASLResponse` message. @@ -487,7 +482,7 @@ mod test { password.as_bytes(), ChannelBinding::unsupported(), nonce.to_string(), - ).unwrap(); + ); assert_eq!(str::from_utf8(scram.message()).unwrap(), client_first); scram.update(server_first.as_bytes()).unwrap(); diff --git a/postgres/src/lib.rs b/postgres/src/lib.rs index eec482bf..d6046015 100644 --- a/postgres/src/lib.rs +++ b/postgres/src/lib.rs @@ -464,7 +464,7 @@ impl InnerConnection { error::connect("a password was requested but not provided".into()) })?; - let mut scram = ScramSha256::new(pass.as_bytes(), channel_binding)?; + let mut scram = ScramSha256::new(pass.as_bytes(), channel_binding); self.stream.write_message(|buf| { frontend::sasl_initial_response(mechanism, scram.message(), buf) @@ -763,8 +763,7 @@ impl InnerConnection { field.name().to_owned(), self.get_type(field.type_oid())?, )) - }) - .collect() + }).collect() .map_err(From::from), None => Ok(vec![]), } @@ -820,7 +819,8 @@ impl InnerConnection { let (name, type_, elem_oid, rngsubtype, basetype, schema, relid) = { let name = String::from_sql_nullable(&Type::NAME, get_raw(0)).map_err(error::conversion)?; - let type_ = i8::from_sql_nullable(&Type::CHAR, get_raw(1)).map_err(error::conversion)?; + let type_ = + i8::from_sql_nullable(&Type::CHAR, get_raw(1)).map_err(error::conversion)?; let elem_oid = Oid::from_sql_nullable(&Type::OID, get_raw(2)).map_err(error::conversion)?; let rngsubtype = Option::::from_sql_nullable(&Type::OID, get_raw(3)) @@ -829,7 +829,8 @@ impl InnerConnection { Oid::from_sql_nullable(&Type::OID, get_raw(4)).map_err(error::conversion)?; let schema = String::from_sql_nullable(&Type::NAME, get_raw(5)).map_err(error::conversion)?; - let relid = Oid::from_sql_nullable(&Type::OID, get_raw(6)).map_err(error::conversion)?; + let relid = + Oid::from_sql_nullable(&Type::OID, get_raw(6)).map_err(error::conversion)?; (name, type_, elem_oid, rngsubtype, basetype, schema, relid) }; @@ -894,7 +895,7 @@ impl InnerConnection { let mut variants = vec![]; for row in rows { variants.push( - String::from_sql_nullable(&Type::NAME, row.get(0)).map_err(error::conversion)? + String::from_sql_nullable(&Type::NAME, row.get(0)).map_err(error::conversion)?, ); } @@ -930,8 +931,8 @@ impl InnerConnection { let mut fields = vec![]; for row in rows { let (name, type_) = { - let name = - String::from_sql_nullable(&Type::NAME, row.get(0)).map_err(error::conversion)?; + let name = String::from_sql_nullable(&Type::NAME, row.get(0)) + .map_err(error::conversion)?; let type_ = Oid::from_sql_nullable(&Type::OID, row.get(1)).map_err(error::conversion)?; (name, type_) diff --git a/tokio-postgres/src/error/mod.rs b/tokio-postgres/src/error/mod.rs index 594f1382..de35c9b1 100644 --- a/tokio-postgres/src/error/mod.rs +++ b/tokio-postgres/src/error/mod.rs @@ -2,10 +2,10 @@ use fallible_iterator::FallibleIterator; use postgres_protocol::message::backend::{ErrorFields, ErrorResponseBody}; -use std::convert::From; use std::error; use std::fmt; use std::io; +use tokio_timer; pub use self::sqlstate::*; @@ -333,147 +333,175 @@ pub enum ErrorPosition { }, } -#[doc(hidden)] -pub fn connect(e: Box) -> Error { - Error(Box::new(ErrorKind::ConnectParams(e))) +#[derive(Debug, PartialEq)] +enum Kind { + Io, + UnexpectedMessage, + Tls, + ToSql, + FromSql, + CopyInStream, + Closed, + Db, + Parse, + Encode, + MissingUser, + MissingPassword, + UnsupportedAuthentication, + Connect, + Timer, + Authentication, } -#[doc(hidden)] -pub fn tls(e: Box) -> Error { - Error(Box::new(ErrorKind::Tls(e))) -} - -#[doc(hidden)] -pub fn db(e: DbError) -> Error { - Error(Box::new(ErrorKind::Db(e))) -} - -#[doc(hidden)] -pub fn __db(e: ErrorResponseBody) -> Error { - match DbError::new(&mut e.fields()) { - Ok(e) => Error(Box::new(ErrorKind::Db(e))), - Err(e) => Error(Box::new(ErrorKind::Io(e))), - } -} - -#[doc(hidden)] -pub fn __user(e: T) -> Error -where - T: Into>, -{ - Error(Box::new(ErrorKind::Conversion(e.into()))) -} - -#[doc(hidden)] -pub fn io(e: io::Error) -> Error { - Error(Box::new(ErrorKind::Io(e))) -} - -#[doc(hidden)] -pub fn conversion(e: Box) -> Error { - Error(Box::new(ErrorKind::Conversion(e))) -} - -#[derive(Debug)] -enum ErrorKind { - ConnectParams(Box), - Tls(Box), - Db(DbError), - Io(io::Error), - Conversion(Box), +struct ErrorInner { + kind: Kind, + cause: Option>, } /// An error communicating with the Postgres server. -#[derive(Debug)] -pub struct Error(Box); +pub struct Error(Box); + +impl fmt::Debug for Error { + fn fmt(&self, fmt: &mut fmt::Formatter) -> fmt::Result { + fmt.debug_struct("Error") + .field("kind", &self.0.kind) + .field("cause", &self.0.cause) + .finish() + } +} impl fmt::Display for Error { fn fmt(&self, fmt: &mut fmt::Formatter) -> fmt::Result { fmt.write_str(error::Error::description(self))?; - match *self.0 { - ErrorKind::ConnectParams(ref err) => write!(fmt, ": {}", err), - ErrorKind::Tls(ref err) => write!(fmt, ": {}", err), - ErrorKind::Db(ref err) => write!(fmt, ": {}", err), - ErrorKind::Io(ref err) => write!(fmt, ": {}", err), - ErrorKind::Conversion(ref err) => write!(fmt, ": {}", err), + if let Some(ref cause) = self.0.cause { + write!(fmt, ": {}", cause)?; } + Ok(()) } } impl error::Error for Error { fn description(&self) -> &str { - match *self.0 { - ErrorKind::ConnectParams(_) => "invalid connection parameters", - ErrorKind::Tls(_) => "TLS handshake error", - ErrorKind::Db(_) => "database error", - ErrorKind::Io(_) => "IO error", - ErrorKind::Conversion(_) => "type conversion error", + match self.0.kind { + Kind::Io => "error communicating with the server", + Kind::UnexpectedMessage => "unexpected message from server", + Kind::Tls => "error performing TLS handshake", + Kind::ToSql => "error serializing a value", + Kind::FromSql => "error deserializing a value", + Kind::CopyInStream => "error from a copy_in stream", + Kind::Closed => "connection closed", + Kind::Db => "db error", + Kind::Parse => "error parsing response from server", + Kind::Encode => "error encoding message to server", + Kind::MissingUser => "username not provided", + Kind::MissingPassword => "password not provided", + Kind::UnsupportedAuthentication => "unsupported authentication method requested", + Kind::Connect => "error connecting to server", + Kind::Timer => "timer error", + Kind::Authentication => "authentication error", } } fn cause(&self) -> Option<&error::Error> { - match *self.0 { - ErrorKind::ConnectParams(ref err) => Some(&**err), - ErrorKind::Tls(ref err) => Some(&**err), - ErrorKind::Db(ref err) => Some(err), - ErrorKind::Io(ref err) => Some(err), - ErrorKind::Conversion(ref err) => Some(&**err), - } + self.0.cause.as_ref().map(|e| &**e as &error::Error) } } impl Error { - /// Returns the SQLSTATE error code associated with this error if it is a DB - /// error. + /// Returns the error's cause. + /// + /// This is the same as `Error::cause` except that it provides extra bounds + /// required to be able to downcast the error. + pub fn cause2(&self) -> Option<&(error::Error + 'static + Sync + Send)> { + self.0.cause.as_ref().map(|e| &**e) + } + + /// Consumes the error, returning its cause. + pub fn into_cause(self) -> Option> { + self.0.cause + } + + /// Returns the SQLSTATE error code associated with the error. + /// + /// This is a convenience method that downcasts the cause to a `DbError` + /// and returns its code. pub fn code(&self) -> Option<&SqlState> { - self.as_db().map(|e| &e.code) + self.cause2() + .and_then(|e| e.downcast_ref::()) + .map(|e| e.code()) } - /// Returns the inner error if this is a connection parameter error. - pub fn as_connection(&self) -> Option<&(error::Error + 'static + Sync + Send)> { - match *self.0 { - ErrorKind::ConnectParams(ref err) => Some(&**err), - _ => None, + fn new(kind: Kind, cause: Option>) -> Error { + Error(Box::new(ErrorInner { kind, cause })) + } + + pub(crate) fn closed() -> Error { + Error::new(Kind::Closed, None) + } + + pub(crate) fn unexpected_message() -> Error { + Error::new(Kind::UnexpectedMessage, None) + } + + pub(crate) fn db(error: ErrorResponseBody) -> Error { + match DbError::new(&mut error.fields()) { + Ok(e) => Error::new(Kind::Db, Some(Box::new(e))), + Err(e) => Error::new(Kind::Parse, Some(Box::new(e))), } } - /// Returns the `DbError` associated with this error if it is a DB error. - pub fn as_db(&self) -> Option<&DbError> { - match *self.0 { - ErrorKind::Db(ref err) => Some(err), - _ => None, - } + pub(crate) fn parse(e: io::Error) -> Error { + Error::new(Kind::Parse, Some(Box::new(e))) } - /// Returns the inner error if this is a conversion error. - pub fn as_conversion(&self) -> Option<&(error::Error + 'static + Sync + Send)> { - match *self.0 { - ErrorKind::Conversion(ref err) => Some(&**err), - _ => None, - } + pub(crate) fn encode(e: io::Error) -> Error { + Error::new(Kind::Encode, Some(Box::new(e))) } - /// Returns the inner `io::Error` associated with this error if it is an IO - /// error. - pub fn as_io(&self) -> Option<&io::Error> { - match *self.0 { - ErrorKind::Io(ref err) => Some(err), - _ => None, - } - } -} - -impl From for Error { - fn from(err: io::Error) -> Error { - Error(Box::new(ErrorKind::Io(err))) - } -} - -impl From for io::Error { - fn from(err: Error) -> io::Error { - match *err.0 { - ErrorKind::Io(e) => e, - _ => io::Error::new(io::ErrorKind::Other, err), - } + pub(crate) fn to_sql(e: Box) -> Error { + Error::new(Kind::ToSql, Some(e)) + } + + pub(crate) fn from_sql(e: Box) -> Error { + Error::new(Kind::FromSql, Some(e)) + } + + pub(crate) fn copy_in_stream(e: E) -> Error + where + E: Into>, + { + Error::new(Kind::CopyInStream, Some(e.into())) + } + + pub(crate) fn missing_user() -> Error { + Error::new(Kind::MissingUser, None) + } + + pub(crate) fn missing_password() -> Error { + Error::new(Kind::MissingPassword, None) + } + + pub(crate) fn unsupported_authentication() -> Error { + Error::new(Kind::UnsupportedAuthentication, None) + } + + pub(crate) fn tls(e: Box) -> Error { + Error::new(Kind::Tls, Some(e)) + } + + pub(crate) fn connect(e: io::Error) -> Error { + Error::new(Kind::Connect, Some(Box::new(e))) + } + + pub(crate) fn timer(e: tokio_timer::Error) -> Error { + Error::new(Kind::Timer, Some(Box::new(e))) + } + + pub(crate) fn io(e: io::Error) -> Error { + Error::new(Kind::Io, Some(Box::new(e))) + } + + pub(crate) fn authentication(e: io::Error) -> Error { + Error::new(Kind::Authentication, Some(Box::new(e))) } } diff --git a/tokio-postgres/src/lib.rs b/tokio-postgres/src/lib.rs index 2ad72bb0..1b8c0b2f 100644 --- a/tokio-postgres/src/lib.rs +++ b/tokio-postgres/src/lib.rs @@ -27,7 +27,6 @@ use futures::{Async, Future, Poll, Stream}; use postgres_shared::rows::RowIndex; use std::error::Error as StdError; use std::fmt; -use std::io; use std::sync::atomic::{AtomicUsize, Ordering}; #[doc(inline)] @@ -52,20 +51,6 @@ fn next_statement() -> String { format!("s{}", NEXT_STATEMENT_ID.fetch_add(1, Ordering::SeqCst)) } -fn bad_response() -> Error { - Error::from(io::Error::new( - io::ErrorKind::InvalidInput, - "the server returned an unexpected response", - )) -} - -fn disconnected() -> Error { - Error::from(io::Error::new( - io::ErrorKind::UnexpectedEof, - "server disconnected", - )) -} - pub enum TlsMode { None, Prefer(Box), diff --git a/tokio-postgres/src/proto/cancel.rs b/tokio-postgres/src/proto/cancel.rs index 21bed4d8..138fb9bb 100644 --- a/tokio-postgres/src/proto/cancel.rs +++ b/tokio-postgres/src/proto/cancel.rs @@ -47,7 +47,7 @@ impl PollCancel for Cancel { fn poll_sending_cancel<'a>( state: &'a mut RentToOwn<'a, SendingCancel>, ) -> Poll { - let (stream, _) = try_ready!(state.future.poll()); + let (stream, _) = try_ready_closed!(state.future.poll()); transition!(FlushingCancel { future: io::flush(stream), @@ -57,7 +57,7 @@ impl PollCancel for Cancel { fn poll_flushing_cancel<'a>( state: &'a mut RentToOwn<'a, FlushingCancel>, ) -> Poll { - try_ready!(state.future.poll()); + try_ready_closed!(state.future.poll()); transition!(Finished(())) } } diff --git a/tokio-postgres/src/proto/client.rs b/tokio-postgres/src/proto/client.rs index eca8956b..4a673305 100644 --- a/tokio-postgres/src/proto/client.rs +++ b/tokio-postgres/src/proto/client.rs @@ -8,8 +8,6 @@ use std::collections::HashMap; use std::error::Error as StdError; use std::sync::{Arc, Weak}; -use disconnected; -use error::{self, Error}; use proto::connection::{Request, RequestMessages}; use proto::copy_in::{CopyInFuture, CopyInReceiver, CopyMessage}; use proto::copy_out::CopyOutStream; @@ -19,6 +17,7 @@ use proto::query::QueryStream; use proto::simple_query::SimpleQueryFuture; use proto::statement::Statement; use types::{IsNull, Oid, ToSql, Type}; +use Error; pub struct PendingRequest(Result); @@ -101,12 +100,12 @@ impl Client { .sender .unbounded_send(Request { messages, sender }) .map(|_| receiver) - .map_err(|_| disconnected()) + .map_err(|_| Error::closed()) } pub fn batch_execute(&self, query: &str) -> SimpleQueryFuture { let pending = self.pending(|buf| { - frontend::query(query, buf)?; + frontend::query(query, buf).map_err(Error::parse)?; Ok(()) }); @@ -115,8 +114,9 @@ impl Client { pub fn prepare(&self, name: String, query: &str, param_types: &[Type]) -> PrepareFuture { let pending = self.pending(|buf| { - frontend::parse(&name, query, param_types.iter().map(|t| t.oid()), buf)?; - frontend::describe(b'S', &name, buf)?; + frontend::parse(&name, query, param_types.iter().map(|t| t.oid()), buf) + .map_err(Error::parse)?; + frontend::describe(b'S', &name, buf).map_err(Error::parse)?; frontend::sync(buf); Ok(()) }); @@ -196,10 +196,10 @@ impl Client { ); match r { Ok(()) => {} - Err(frontend::BindError::Conversion(e)) => return Err(error::conversion(e)), - Err(frontend::BindError::Serialization(e)) => return Err(Error::from(e)), + Err(frontend::BindError::Conversion(e)) => return Err(Error::to_sql(e)), + Err(frontend::BindError::Serialization(e)) => return Err(Error::encode(e)), } - frontend::execute("", 0, &mut buf)?; + frontend::execute("", 0, &mut buf).map_err(Error::parse)?; frontend::sync(&mut buf); Ok(buf) } diff --git a/tokio-postgres/src/proto/connect.rs b/tokio-postgres/src/proto/connect.rs index 5d7b175f..96c0a0ce 100644 --- a/tokio-postgres/src/proto/connect.rs +++ b/tokio-postgres/src/proto/connect.rs @@ -14,11 +14,10 @@ use tokio_timer::Delay; #[cfg(unix)] use tokio_uds::{self, UnixStream}; -use error::{self, Error}; use params::{ConnectParams, Host}; use proto::socket::Socket; use tls::{self, TlsConnect, TlsStream}; -use {bad_response, TlsMode}; +use {Error, TlsMode}; lazy_static! { static ref DNS_POOL: CpuPool = CpuPool::new(2); @@ -27,7 +26,10 @@ lazy_static! { #[derive(StateMachineFuture)] pub enum Connect { #[state_machine_future(start)] - #[cfg_attr(unix, state_machine_future(transitions(ResolvingDns, ConnectingUnix)))] + #[cfg_attr( + unix, + state_machine_future(transitions(ResolvingDns, ConnectingUnix)) + )] #[cfg_attr(not(unix), state_machine_future(transitions(ResolvingDns)))] Start { params: ConnectParams, tls: TlsMode }, #[state_machine_future(transitions(ConnectingTcp))] @@ -122,13 +124,16 @@ impl PollConnect for Connect { fn poll_resolving_dns<'a>( state: &'a mut RentToOwn<'a, ResolvingDns>, ) -> Poll { - let mut addrs = try_ready!(state.future.poll()); + let mut addrs = try_ready!(state.future.poll().map_err(Error::connect)); let state = state.take(); let addr = match addrs.next() { Some(addr) => addr, None => { - return Err(io::Error::new(io::ErrorKind::Other, "resolved to 0 addresses").into()) + return Err(Error::connect(io::Error::new( + io::ErrorKind::Other, + "resolved to 0 addresses", + ))) } }; @@ -149,11 +154,7 @@ impl PollConnect for Connect { Ok(Async::Ready(socket)) => break socket, Ok(Async::NotReady) => match state.timeout { Some((_, ref mut delay)) => { - try_ready!( - delay - .poll() - .map_err(|e| io::Error::new(io::ErrorKind::Other, e)) - ); + try_ready!(delay.poll().map_err(Error::timer)); io::Error::new(io::ErrorKind::TimedOut, "connection timed out") } None => return Ok(Async::NotReady), @@ -163,7 +164,7 @@ impl PollConnect for Connect { let addr = match state.addrs.next() { Some(addr) => addr, - None => return Err(error.into()), + None => return Err(Error::connect(error)), }; state.future = TcpStream::connect(&addr); @@ -175,7 +176,7 @@ impl PollConnect for Connect { // Our read/write patterns may trigger Nagle's algorithm since we're pipelining which // we don't want. Each individual write should be a full command we want the backend to // see immediately. - socket.set_nodelay(true)?; + socket.set_nodelay(true).map_err(Error::connect)?; let state = state.take(); transition!(PreparingSsl { @@ -189,7 +190,7 @@ impl PollConnect for Connect { fn poll_connecting_unix<'a>( state: &'a mut RentToOwn<'a, ConnectingUnix>, ) -> Poll { - match state.future.poll()? { + match state.future.poll().map_err(Error::connect)? { Async::Ready(socket) => { let state = state.take(); transition!(PreparingSsl { @@ -200,12 +201,11 @@ impl PollConnect for Connect { } Async::NotReady => match state.timeout { Some(ref mut delay) => { - try_ready!( - delay - .poll() - .map_err(|e| io::Error::new(io::ErrorKind::Other, e)) - ); - Err(io::Error::new(io::ErrorKind::TimedOut, "connection timed out").into()) + try_ready!(delay.poll().map_err(Error::timer)); + Err(Error::connect(io::Error::new( + io::ErrorKind::TimedOut, + "connection timed out", + ))) } None => Ok(Async::NotReady), }, @@ -238,7 +238,7 @@ impl PollConnect for Connect { fn poll_sending_ssl<'a>( state: &'a mut RentToOwn<'a, SendingSsl>, ) -> Poll { - let (stream, _) = try_ready!(state.future.poll()); + let (stream, _) = try_ready_closed!(state.future.poll()); let state = state.take(); transition!(FlushingSsl { future: flush(stream), @@ -251,7 +251,7 @@ impl PollConnect for Connect { fn poll_flushing_ssl<'a>( state: &'a mut RentToOwn<'a, FlushingSsl>, ) -> Poll { - let stream = try_ready!(state.future.poll()); + let stream = try_ready_closed!(state.future.poll()); let state = state.take(); transition!(ReadingSsl { future: read_exact(stream, [0]), @@ -264,7 +264,7 @@ impl PollConnect for Connect { fn poll_reading_ssl<'a>( state: &'a mut RentToOwn<'a, ReadingSsl>, ) -> Poll { - let (stream, buf) = try_ready!(state.future.poll()); + let (stream, buf) = try_ready_closed!(state.future.poll()); let state = state.take(); match buf[0] { @@ -272,7 +272,7 @@ impl PollConnect for Connect { let future = match state.params.host() { Host::Tcp(domain) => state.connector.connect(domain, tls::Socket(stream)), Host::Unix(_) => { - return Err(error::tls("TLS over unix sockets not supported".into())) + return Err(Error::tls("TLS over unix sockets not supported".into())) } }; transition!(ConnectingTls { @@ -281,15 +281,15 @@ impl PollConnect for Connect { }) } b'N' if !state.required => transition!(Ready(Box::new(stream))), - b'N' => Err(error::tls("TLS was required but not supported".into())), - _ => Err(bad_response()), + b'N' => Err(Error::tls("TLS was required but not supported".into())), + _ => Err(Error::unexpected_message()), } } fn poll_connecting_tls<'a>( state: &'a mut RentToOwn<'a, ConnectingTls>, ) -> Poll { - let stream = try_ready!(state.future.poll().map_err(error::tls)); + let stream = try_ready!(state.future.poll().map_err(Error::tls)); transition!(Ready(stream)) } } diff --git a/tokio-postgres/src/proto/connection.rs b/tokio-postgres/src/proto/connection.rs index a2b01c8d..562ec6ee 100644 --- a/tokio-postgres/src/proto/connection.rs +++ b/tokio-postgres/src/proto/connection.rs @@ -6,11 +6,11 @@ use std::collections::{HashMap, VecDeque}; use std::io; use tokio_codec::Framed; -use error::{self, DbError, Error}; use proto::codec::PostgresCodec; use proto::copy_in::CopyInReceiver; use tls::TlsStream; -use {bad_response, disconnected, AsyncMessage, CancelData, Notification}; +use {AsyncMessage, CancelData, Notification}; +use {DbError, Error}; pub enum RequestMessages { Single(Vec), @@ -86,10 +86,10 @@ impl Connection { } loop { - let message = match self.poll_response()? { + let message = match self.poll_response().map_err(Error::io)? { Async::Ready(Some(message)) => message, Async::Ready(None) => { - return Err(disconnected()); + return Err(Error::closed()); } Async::NotReady => { trace!("poll_read: waiting on response"); @@ -99,20 +99,22 @@ impl Connection { let message = match message { Message::NoticeResponse(body) => { - let error = DbError::new(&mut body.fields())?; + let error = DbError::new(&mut body.fields()).map_err(Error::parse)?; return Ok(Some(AsyncMessage::Notice(error))); } Message::NotificationResponse(body) => { let notification = Notification { process_id: body.process_id(), - channel: body.channel()?.to_string(), - payload: body.message()?.to_string(), + channel: body.channel().map_err(Error::parse)?.to_string(), + payload: body.message().map_err(Error::parse)?.to_string(), }; return Ok(Some(AsyncMessage::Notification(notification))); } Message::ParameterStatus(body) => { - self.parameters - .insert(body.name()?.to_string(), body.value()?.to_string()); + self.parameters.insert( + body.name().map_err(Error::parse)?.to_string(), + body.value().map_err(Error::parse)?.to_string(), + ); continue; } m => m, @@ -121,8 +123,8 @@ impl Connection { let mut sender = match self.responses.pop_front() { Some(sender) => sender, None => match message { - Message::ErrorResponse(error) => return Err(error::__db(error)), - _ => return Err(bad_response()), + Message::ErrorResponse(error) => return Err(Error::db(error)), + _ => return Err(Error::unexpected_message()), }, }; @@ -155,7 +157,7 @@ impl Connection { return Ok(Async::Ready(Some(message))); } - match try_receive!(self.receiver.poll()) { + match try_ready_receive!(self.receiver.poll()) { Some(request) => { trace!("polled new request"); self.responses.push_back(request.sender); @@ -195,19 +197,21 @@ impl Connection { }; match request { - RequestMessages::Single(request) => match self.stream.start_send(request)? { - AsyncSink::Ready => { - if self.state == State::Terminating { - trace!("poll_write: sent eof, closing"); - self.state = State::Closing; + RequestMessages::Single(request) => { + match self.stream.start_send(request).map_err(Error::io)? { + AsyncSink::Ready => { + if self.state == State::Terminating { + trace!("poll_write: sent eof, closing"); + self.state = State::Closing; + } + } + AsyncSink::NotReady(request) => { + trace!("poll_write: waiting on socket"); + self.pending_request = Some(RequestMessages::Single(request)); + return Ok(false); } } - AsyncSink::NotReady(request) => { - trace!("poll_write: waiting on socket"); - self.pending_request = Some(RequestMessages::Single(request)); - return Ok(false); - } - }, + } RequestMessages::CopyIn { mut receiver, mut pending_message, @@ -232,7 +236,7 @@ impl Connection { }, }; - match self.stream.start_send(message)? { + match self.stream.start_send(message).map_err(Error::io)? { AsyncSink::Ready => { self.pending_request = Some(RequestMessages::CopyIn { receiver, @@ -254,17 +258,11 @@ impl Connection { } fn poll_flush(&mut self) -> Result<(), Error> { - match self.stream.poll_complete() { - Ok(Async::Ready(())) => { - trace!("poll_flush: flushed"); - Ok(()) - } - Ok(Async::NotReady) => { - trace!("poll_flush: waiting on socket"); - Ok(()) - } - Err(e) => Err(Error::from(e)), + match self.stream.poll_complete().map_err(Error::io)? { + Async::Ready(()) => trace!("poll_flush: flushed"), + Async::NotReady => trace!("poll_flush: waiting on socket"), } + Ok(()) } fn poll_shutdown(&mut self) -> Poll<(), Error> { @@ -272,16 +270,15 @@ impl Connection { return Ok(Async::NotReady); } - match self.stream.close() { - Ok(Async::Ready(())) => { + match self.stream.close().map_err(Error::io)? { + Async::Ready(()) => { trace!("poll_shutdown: complete"); Ok(Async::Ready(())) } - Ok(Async::NotReady) => { + Async::NotReady => { trace!("poll_shutdown: waiting on socket"); Ok(Async::NotReady) } - Err(e) => Err(Error::from(e)), } } diff --git a/tokio-postgres/src/proto/copy_in.rs b/tokio-postgres/src/proto/copy_in.rs index cec324d2..8f22dbca 100644 --- a/tokio-postgres/src/proto/copy_in.rs +++ b/tokio-postgres/src/proto/copy_in.rs @@ -6,10 +6,9 @@ use postgres_protocol::message::frontend; use state_machine_future::RentToOwn; use std::error::Error as StdError; -use error::{self, Error}; use proto::client::{Client, PendingRequest}; use proto::statement::Statement; -use {bad_response, disconnected}; +use Error; pub enum CopyMessage { Data(Vec), @@ -123,7 +122,7 @@ where state: &'a mut RentToOwn<'a, ReadCopyInResponse>, ) -> Poll, Error> { loop { - let message = try_receive!(state.receiver.poll()); + let message = try_ready_receive!(state.receiver.poll()); match message { Some(Message::BindComplete) => {} @@ -136,9 +135,9 @@ where receiver: state.receiver }) } - Some(Message::ErrorResponse(body)) => return Err(error::__db(body)), - Some(_) => return Err(bad_response()), - None => return Err(disconnected()), + Some(Message::ErrorResponse(body)) => return Err(Error::db(body)), + Some(_) => return Err(Error::unexpected_message()), + None => return Err(Error::closed()), } } } @@ -149,10 +148,10 @@ where loop { let message = match state.pending_message.take() { Some(message) => message, - None => match try_ready!(state.stream.poll().map_err(error::__user)) { + None => match try_ready!(state.stream.poll().map_err(Error::copy_in_stream)) { Some(data) => { let mut buf = vec![]; - frontend::copy_data(data.as_ref(), &mut buf).map_err(error::io)?; + frontend::copy_data(data.as_ref(), &mut buf).map_err(Error::encode)?; CopyMessage::Data(buf) } None => { @@ -171,7 +170,7 @@ where state.pending_message = Some(message); return Ok(Async::NotReady); } - Err(_) => return Err(disconnected()), + Err(_) => return Err(Error::closed()), } } } @@ -179,7 +178,7 @@ where fn poll_write_copy_done<'a>( state: &'a mut RentToOwn<'a, WriteCopyDone>, ) -> Poll { - try_ready!(state.future.poll().map_err(|_| disconnected())); + try_ready!(state.future.poll().map_err(|_| Error::closed())); let state = state.take(); transition!(ReadCommandComplete { @@ -190,16 +189,23 @@ where fn poll_read_command_complete<'a>( state: &'a mut RentToOwn<'a, ReadCommandComplete>, ) -> Poll { - let message = try_receive!(state.receiver.poll()); + let message = try_ready_receive!(state.receiver.poll()); match message { Some(Message::CommandComplete(body)) => { - let rows = body.tag()?.rsplit(' ').next().unwrap().parse().unwrap_or(0); + let rows = body + .tag() + .map_err(Error::parse)? + .rsplit(' ') + .next() + .unwrap() + .parse() + .unwrap_or(0); transition!(Finished(rows)) } - Some(Message::ErrorResponse(body)) => Err(error::__db(body)), - Some(_) => Err(bad_response()), - None => Err(disconnected()), + Some(Message::ErrorResponse(body)) => Err(Error::db(body)), + Some(_) => Err(Error::unexpected_message()), + None => Err(Error::closed()), } } } diff --git a/tokio-postgres/src/proto/copy_out.rs b/tokio-postgres/src/proto/copy_out.rs index b8909d5d..2fdf1dbf 100644 --- a/tokio-postgres/src/proto/copy_out.rs +++ b/tokio-postgres/src/proto/copy_out.rs @@ -4,10 +4,9 @@ use futures::{Async, Poll, Stream}; use postgres_protocol::message::backend::Message; use std::mem; -use error::{self, Error}; use proto::client::{Client, PendingRequest}; use proto::statement::Statement; -use {bad_response, disconnected}; +use Error; enum State { Start { @@ -60,9 +59,9 @@ impl Stream for CopyOutStream { Some(Message::CopyOutResponse(_)) => { self.0 = State::ReadingCopyData { receiver }; } - Some(Message::ErrorResponse(body)) => break Err(error::__db(body)), - Some(_) => break Err(bad_response()), - None => break Err(disconnected()), + Some(Message::ErrorResponse(body)) => break Err(Error::db(body)), + Some(_) => break Err(Error::unexpected_message()), + None => break Err(Error::closed()), } } State::ReadingCopyData { mut receiver } => { @@ -84,9 +83,9 @@ impl Stream for CopyOutStream { self.0 = State::ReadingCopyData { receiver }; } Some(Message::ReadyForQuery(_)) => break Ok(Async::Ready(None)), - Some(Message::ErrorResponse(body)) => break Err(error::__db(body)), - Some(_) => break Err(bad_response()), - None => break Err(disconnected()), + Some(Message::ErrorResponse(body)) => break Err(Error::db(body)), + Some(_) => break Err(Error::unexpected_message()), + None => break Err(Error::closed()), } } State::Done => break Ok(Async::Ready(None)), diff --git a/tokio-postgres/src/proto/execute.rs b/tokio-postgres/src/proto/execute.rs index 67315c88..34b08c20 100644 --- a/tokio-postgres/src/proto/execute.rs +++ b/tokio-postgres/src/proto/execute.rs @@ -3,10 +3,9 @@ use futures::{Poll, Stream}; use postgres_protocol::message::backend::Message; use state_machine_future::RentToOwn; -use error::{self, Error}; use proto::client::{Client, PendingRequest}; use proto::statement::Statement; -use {bad_response, disconnected}; +use Error; #[derive(StateMachineFuture)] pub enum Execute { @@ -42,14 +41,21 @@ impl PollExecute for Execute { state: &'a mut RentToOwn<'a, ReadResponse>, ) -> Poll { loop { - let message = try_receive!(state.receiver.poll()); + let message = try_ready_receive!(state.receiver.poll()); match message { Some(Message::BindComplete) => {} Some(Message::DataRow(_)) => {} - Some(Message::ErrorResponse(body)) => return Err(error::__db(body)), + Some(Message::ErrorResponse(body)) => return Err(Error::db(body)), Some(Message::CommandComplete(body)) => { - let rows = body.tag()?.rsplit(' ').next().unwrap().parse().unwrap_or(0); + let rows = body + .tag() + .map_err(Error::parse)? + .rsplit(' ') + .next() + .unwrap() + .parse() + .unwrap_or(0); let state = state.take(); transition!(ReadReadyForQuery { receiver: state.receiver, @@ -63,8 +69,8 @@ impl PollExecute for Execute { rows: 0, }); } - Some(_) => return Err(bad_response()), - None => return Err(disconnected()), + Some(_) => return Err(Error::unexpected_message()), + None => return Err(Error::closed()), } } } @@ -72,12 +78,12 @@ impl PollExecute for Execute { fn poll_read_ready_for_query<'a>( state: &'a mut RentToOwn<'a, ReadReadyForQuery>, ) -> Poll { - let message = try_receive!(state.receiver.poll()); + let message = try_ready_receive!(state.receiver.poll()); match message { Some(Message::ReadyForQuery(_)) => transition!(Finished(state.rows)), - Some(_) => Err(bad_response()), - None => Err(disconnected()), + Some(_) => Err(Error::unexpected_message()), + None => Err(Error::closed()), } } } diff --git a/tokio-postgres/src/proto/handshake.rs b/tokio-postgres/src/proto/handshake.rs index 5975c496..119669f7 100644 --- a/tokio-postgres/src/proto/handshake.rs +++ b/tokio-postgres/src/proto/handshake.rs @@ -11,14 +11,13 @@ use std::collections::HashMap; use std::io; use tokio_codec::Framed; -use error::{self, Error}; use params::{ConnectParams, User}; use proto::client::Client; use proto::codec::PostgresCodec; use proto::connect::ConnectFuture; use proto::connection::Connection; use tls::TlsStream; -use {bad_response, disconnected, CancelData, TlsMode}; +use {CancelData, Error, TlsMode}; #[derive(StateMachineFuture)] pub enum Handshake { @@ -74,11 +73,7 @@ impl PollHandshake for Handshake { let user = match state.params.user() { Some(user) => user.clone(), - None => { - return Err(error::connect( - "user missing from connection parameters".into(), - )) - } + None => return Err(Error::missing_user()), }; let mut buf = vec![]; @@ -100,7 +95,7 @@ impl PollHandshake for Handshake { .chain(user) .chain(database), &mut buf, - )?; + ).map_err(Error::encode)?; } let stream = Framed::new(stream, PostgresCodec); @@ -113,7 +108,7 @@ impl PollHandshake for Handshake { fn poll_sending_startup<'a>( state: &'a mut RentToOwn<'a, SendingStartup>, ) -> Poll { - let stream = try_ready!(state.future.poll()); + let stream = try_ready!(state.future.poll().map_err(Error::io)); let state = state.take(); transition!(ReadingAuth { stream, @@ -124,7 +119,7 @@ impl PollHandshake for Handshake { fn poll_reading_auth<'a>( state: &'a mut RentToOwn<'a, ReadingAuth>, ) -> Poll { - let message = try_ready!(state.stream.poll()); + let message = try_ready!(state.stream.poll().map_err(Error::io)); let state = state.take(); match message { @@ -134,33 +129,33 @@ impl PollHandshake for Handshake { parameters: HashMap::new(), }), Some(Message::AuthenticationCleartextPassword) => { - let pass = state.user.password().ok_or_else(missing_password)?; + let pass = state.user.password().ok_or_else(Error::missing_password)?; let mut buf = vec![]; - frontend::password_message(pass, &mut buf)?; + frontend::password_message(pass, &mut buf).map_err(Error::encode)?; transition!(SendingPassword { future: state.stream.send(buf) }) } Some(Message::AuthenticationMd5Password(body)) => { - let pass = state.user.password().ok_or_else(missing_password)?; + let pass = state.user.password().ok_or_else(Error::missing_password)?; let output = authentication::md5_hash( state.user.name().as_bytes(), pass.as_bytes(), body.salt(), ); let mut buf = vec![]; - frontend::password_message(&output, &mut buf)?; + frontend::password_message(&output, &mut buf).map_err(Error::encode)?; transition!(SendingPassword { future: state.stream.send(buf) }) } Some(Message::AuthenticationSasl(body)) => { - let pass = state.user.password().ok_or_else(missing_password)?; + let pass = state.user.password().ok_or_else(Error::missing_password)?; let mut has_scram = false; let mut has_scram_plus = false; let mut mechanisms = body.mechanisms(); - while let Some(mechanism) = mechanisms.next()? { + while let Some(mechanism) = mechanisms.next().map_err(Error::parse)? { match mechanism { sasl::SCRAM_SHA_256 => has_scram = true, sasl::SCRAM_SHA_256_PLUS => has_scram_plus = true, @@ -191,16 +186,14 @@ impl PollHandshake for Handshake { None => (ChannelBinding::unsupported(), sasl::SCRAM_SHA_256), } } else { - return Err(io::Error::new( - io::ErrorKind::Other, - "unsupported SASL authentication", - ).into()); + return Err(Error::unsupported_authentication()); }; - let mut scram = ScramSha256::new(pass.as_bytes(), channel_binding)?; + let mut scram = ScramSha256::new(pass.as_bytes(), channel_binding); let mut buf = vec![]; - frontend::sasl_initial_response(mechanism, scram.message(), &mut buf)?; + frontend::sasl_initial_response(mechanism, scram.message(), &mut buf) + .map_err(Error::encode)?; transition!(SendingSasl { future: state.stream.send(buf), @@ -210,27 +203,24 @@ impl PollHandshake for Handshake { Some(Message::AuthenticationKerberosV5) | Some(Message::AuthenticationScmCredential) | Some(Message::AuthenticationGss) - | Some(Message::AuthenticationSspi) => Err(io::Error::new( - io::ErrorKind::Other, - "unsupported authentication method", - ).into()), - Some(Message::ErrorResponse(body)) => Err(error::__db(body)), - Some(_) => Err(bad_response()), - None => Err(disconnected()), + | Some(Message::AuthenticationSspi) => Err(Error::unsupported_authentication()), + Some(Message::ErrorResponse(body)) => Err(Error::db(body)), + Some(_) => Err(Error::unexpected_message()), + None => Err(Error::closed()), } } fn poll_sending_password<'a>( state: &'a mut RentToOwn<'a, SendingPassword>, ) -> Poll { - let stream = try_ready!(state.future.poll()); + let stream = try_ready!(state.future.poll().map_err(Error::io)); transition!(ReadingAuthCompletion { stream }) } fn poll_sending_sasl<'a>( state: &'a mut RentToOwn<'a, SendingSasl>, ) -> Poll { - let stream = try_ready!(state.future.poll()); + let stream = try_ready!(state.future.poll().map_err(Error::io)); let state = state.take(); transition!(ReadingSasl { stream, @@ -241,35 +231,41 @@ impl PollHandshake for Handshake { fn poll_reading_sasl<'a>( state: &'a mut RentToOwn<'a, ReadingSasl>, ) -> Poll { - let message = try_ready!(state.stream.poll()); + let message = try_ready!(state.stream.poll().map_err(Error::io)); let mut state = state.take(); match message { Some(Message::AuthenticationSaslContinue(body)) => { - state.scram.update(body.data())?; + state + .scram + .update(body.data()) + .map_err(Error::authentication)?; let mut buf = vec![]; - frontend::sasl_response(state.scram.message(), &mut buf)?; + frontend::sasl_response(state.scram.message(), &mut buf).map_err(Error::encode)?; transition!(SendingSasl { future: state.stream.send(buf), scram: state.scram, }) } Some(Message::AuthenticationSaslFinal(body)) => { - state.scram.finish(body.data())?; + state + .scram + .finish(body.data()) + .map_err(Error::authentication)?; transition!(ReadingAuthCompletion { stream: state.stream, }) } - Some(Message::ErrorResponse(body)) => Err(error::__db(body)), - Some(_) => Err(bad_response()), - None => Err(disconnected()), + Some(Message::ErrorResponse(body)) => Err(Error::db(body)), + Some(_) => Err(Error::unexpected_message()), + None => Err(Error::closed()), } } fn poll_reading_auth_completion<'a>( state: &'a mut RentToOwn<'a, ReadingAuthCompletion>, ) -> Poll { - let message = try_ready!(state.stream.poll()); + let message = try_ready!(state.stream.poll().map_err(Error::io)); let state = state.take(); match message { @@ -278,9 +274,9 @@ impl PollHandshake for Handshake { cancel_data: None, parameters: HashMap::new(), }), - Some(Message::ErrorResponse(body)) => Err(error::__db(body)), - Some(_) => Err(bad_response()), - None => Err(disconnected()), + Some(Message::ErrorResponse(body)) => Err(Error::db(body)), + Some(_) => Err(Error::unexpected_message()), + None => Err(Error::closed()), } } @@ -288,7 +284,7 @@ impl PollHandshake for Handshake { state: &'a mut RentToOwn<'a, ReadingInfo>, ) -> Poll { loop { - let message = try_ready!(state.stream.poll()); + let message = try_ready!(state.stream.poll().map_err(Error::io)); match message { Some(Message::BackendKeyData(body)) => { state.cancel_data = Some(CancelData { @@ -297,14 +293,18 @@ impl PollHandshake for Handshake { }); } Some(Message::ParameterStatus(body)) => { - state - .parameters - .insert(body.name()?.to_string(), body.value()?.to_string()); + state.parameters.insert( + body.name().map_err(Error::parse)?.to_string(), + body.value().map_err(Error::parse)?.to_string(), + ); } Some(Message::ReadyForQuery(_)) => { let state = state.take(); let cancel_data = state.cancel_data.ok_or_else(|| { - io::Error::new(io::ErrorKind::InvalidData, "BackendKeyData message missing") + Error::parse(io::Error::new( + io::ErrorKind::InvalidData, + "BackendKeyData message missing", + )) })?; let (sender, receiver) = mpsc::unbounded(); let client = Client::new(sender); @@ -312,10 +312,10 @@ impl PollHandshake for Handshake { Connection::new(state.stream, cancel_data, state.parameters, receiver); transition!(Finished((client, connection))) } - Some(Message::ErrorResponse(body)) => return Err(error::__db(body)), + Some(Message::ErrorResponse(body)) => return Err(Error::db(body)), Some(Message::NoticeResponse(_)) => {} - Some(_) => return Err(bad_response()), - None => return Err(disconnected()), + Some(_) => return Err(Error::unexpected_message()), + None => return Err(Error::closed()), } } } @@ -326,7 +326,3 @@ impl HandshakeFuture { Handshake::start(ConnectFuture::new(params.clone(), tls), params) } } - -fn missing_password() -> Error { - error::connect("a password was requested but not provided".into()) -} diff --git a/tokio-postgres/src/proto/mod.rs b/tokio-postgres/src/proto/mod.rs index f5d81872..1f782773 100644 --- a/tokio-postgres/src/proto/mod.rs +++ b/tokio-postgres/src/proto/mod.rs @@ -1,4 +1,4 @@ -macro_rules! try_receive { +macro_rules! try_ready_receive { ($e:expr) => { match $e { Ok(::futures::Async::Ready(v)) => v, @@ -8,6 +8,16 @@ macro_rules! try_receive { }; } +macro_rules! try_ready_closed { + ($e:expr) => { + match $e { + Ok(::futures::Async::Ready(v)) => v, + Ok(::futures::Async::NotReady) => return Ok(::futures::Async::NotReady), + Err(_) => return Err(::Error::closed()), + } + }; +} + mod cancel; mod client; mod codec; diff --git a/tokio-postgres/src/proto/prepare.rs b/tokio-postgres/src/proto/prepare.rs index 2252189e..8fe6003a 100644 --- a/tokio-postgres/src/proto/prepare.rs +++ b/tokio-postgres/src/proto/prepare.rs @@ -6,13 +6,11 @@ use state_machine_future::RentToOwn; use std::mem; use std::vec; -use error::{self, Error}; use proto::client::{Client, PendingRequest}; use proto::statement::Statement; use proto::typeinfo::TypeinfoFuture; use types::{Oid, Type}; -use Column; -use {bad_response, disconnected}; +use {Column, Error}; #[derive(StateMachineFuture)] pub enum Prepare { @@ -87,7 +85,7 @@ impl PollPrepare for Prepare { fn poll_read_parse_complete<'a>( state: &'a mut RentToOwn<'a, ReadParseComplete>, ) -> Poll { - let message = try_receive!(state.receiver.poll()); + let message = try_ready_receive!(state.receiver.poll()); let state = state.take(); match message { @@ -96,44 +94,45 @@ impl PollPrepare for Prepare { name: state.name, client: state.client, }), - Some(Message::ErrorResponse(body)) => Err(error::__db(body)), - Some(_) => Err(bad_response()), - None => Err(disconnected()), + Some(Message::ErrorResponse(body)) => Err(Error::db(body)), + Some(_) => Err(Error::unexpected_message()), + None => Err(Error::closed()), } } fn poll_read_parameter_description<'a>( state: &'a mut RentToOwn<'a, ReadParameterDescription>, ) -> Poll { - let message = try_receive!(state.receiver.poll()); + let message = try_ready_receive!(state.receiver.poll()); let state = state.take(); match message { Some(Message::ParameterDescription(body)) => transition!(ReadRowDescription { receiver: state.receiver, name: state.name, - parameters: body.parameters().collect()?, + parameters: body.parameters().collect().map_err(Error::parse)?, client: state.client, }), - Some(_) => Err(bad_response()), - None => Err(disconnected()), + Some(_) => Err(Error::unexpected_message()), + None => Err(Error::closed()), } } fn poll_read_row_description<'a>( state: &'a mut RentToOwn<'a, ReadRowDescription>, ) -> Poll { - let message = try_receive!(state.receiver.poll()); + let message = try_ready_receive!(state.receiver.poll()); let state = state.take(); let columns = match message { Some(Message::RowDescription(body)) => body .fields() .map(|f| (f.name().to_string(), f.type_oid())) - .collect()?, + .collect() + .map_err(Error::parse)?, Some(Message::NoData) => vec![], - Some(_) => return Err(bad_response()), - None => return Err(disconnected()), + Some(_) => return Err(Error::unexpected_message()), + None => return Err(Error::closed()), }; transition!(ReadReadyForQuery { @@ -148,13 +147,13 @@ impl PollPrepare for Prepare { fn poll_read_ready_for_query<'a>( state: &'a mut RentToOwn<'a, ReadReadyForQuery>, ) -> Poll { - let message = try_receive!(state.receiver.poll()); + let message = try_ready_receive!(state.receiver.poll()); let state = state.take(); match message { Some(Message::ReadyForQuery(_)) => {} - Some(_) => return Err(bad_response()), - None => return Err(disconnected()), + Some(_) => return Err(Error::unexpected_message()), + None => return Err(Error::closed()), } let mut parameters = state.parameters.into_iter(); diff --git a/tokio-postgres/src/proto/query.rs b/tokio-postgres/src/proto/query.rs index 5922b481..985bdca3 100644 --- a/tokio-postgres/src/proto/query.rs +++ b/tokio-postgres/src/proto/query.rs @@ -3,11 +3,10 @@ use futures::{Async, Poll, Stream}; use postgres_protocol::message::backend::Message; use std::mem; -use error::{self, Error}; use proto::client::{Client, PendingRequest}; use proto::row::Row; use proto::statement::Statement; -use {bad_response, disconnected}; +use Error; enum State { Start { @@ -68,7 +67,7 @@ impl Stream for QueryStream { statement, }; } - Some(Message::ErrorResponse(body)) => break Err(error::__db(body)), + Some(Message::ErrorResponse(body)) => break Err(Error::db(body)), Some(Message::DataRow(body)) => { let row = Row::new(statement.clone(), body)?; self.0 = State::ReadingResponse { @@ -80,8 +79,8 @@ impl Stream for QueryStream { Some(Message::EmptyQueryResponse) | Some(Message::CommandComplete(_)) => { self.0 = State::ReadingReadyForQuery { receiver }; } - Some(_) => break Err(bad_response()), - None => break Err(disconnected()), + Some(_) => break Err(Error::unexpected_message()), + None => break Err(Error::closed()), } } State::ReadingReadyForQuery { mut receiver } => { @@ -96,8 +95,8 @@ impl Stream for QueryStream { match message { Some(Message::ReadyForQuery(_)) => break Ok(Async::Ready(None)), - Some(_) => break Err(bad_response()), - None => break Err(disconnected()), + Some(_) => break Err(Error::unexpected_message()), + None => break Err(Error::closed()), } } State::Done => break Ok(Async::Ready(None)), diff --git a/tokio-postgres/src/proto/row.rs b/tokio-postgres/src/proto/row.rs index 38897410..38c270d4 100644 --- a/tokio-postgres/src/proto/row.rs +++ b/tokio-postgres/src/proto/row.rs @@ -2,10 +2,9 @@ use postgres_protocol::message::backend::DataRowBody; use postgres_shared::rows::{RowData, RowIndex}; use std::fmt; -use error::{self, Error}; use proto::statement::Statement; use types::{FromSql, WrongType}; -use Column; +use {Column, Error}; pub struct Row { statement: Statement, @@ -14,7 +13,7 @@ pub struct Row { impl Row { pub fn new(statement: Statement, data: DataRowBody) -> Result { - let data = RowData::new(data)?; + let data = RowData::new(data).map_err(Error::parse)?; Ok(Row { statement, data }) } @@ -58,9 +57,9 @@ impl Row { let ty = self.statement.columns()[idx].type_(); if !::accepts(ty) { - return Err(error::conversion(Box::new(WrongType::new(ty.clone())))); + return Err(Error::from_sql(Box::new(WrongType::new(ty.clone())))); } let value = FromSql::from_sql_nullable(ty, self.data.get(idx)); - value.map(Some).map_err(error::conversion) + value.map(Some).map_err(Error::from_sql) } } diff --git a/tokio-postgres/src/proto/simple_query.rs b/tokio-postgres/src/proto/simple_query.rs index b3a4d75c..e39d1b4e 100644 --- a/tokio-postgres/src/proto/simple_query.rs +++ b/tokio-postgres/src/proto/simple_query.rs @@ -3,9 +3,8 @@ use futures::{Poll, Stream}; use postgres_protocol::message::backend::Message; use state_machine_future::RentToOwn; -use error::{self, Error}; use proto::client::{Client, PendingRequest}; -use {bad_response, disconnected}; +use Error; #[derive(StateMachineFuture)] pub enum SimpleQuery { @@ -34,17 +33,17 @@ impl PollSimpleQuery for SimpleQuery { state: &'a mut RentToOwn<'a, ReadResponse>, ) -> Poll { loop { - let message = try_receive!(state.receiver.poll()); + let message = try_ready_receive!(state.receiver.poll()); match message { Some(Message::CommandComplete(_)) | Some(Message::RowDescription(_)) | Some(Message::DataRow(_)) | Some(Message::EmptyQueryResponse) => {} - Some(Message::ErrorResponse(body)) => return Err(error::__db(body)), + Some(Message::ErrorResponse(body)) => return Err(Error::db(body)), Some(Message::ReadyForQuery(_)) => transition!(Finished(())), - Some(_) => return Err(bad_response()), - None => return Err(disconnected()), + Some(_) => return Err(Error::unexpected_message()), + None => return Err(Error::closed()), } } } diff --git a/tokio-postgres/src/proto/typeinfo.rs b/tokio-postgres/src/proto/typeinfo.rs index f4900e0c..86d75f0d 100644 --- a/tokio-postgres/src/proto/typeinfo.rs +++ b/tokio-postgres/src/proto/typeinfo.rs @@ -3,13 +3,13 @@ use futures::{Async, Future, Poll}; use state_machine_future::RentToOwn; use error::{Error, SqlState}; +use next_statement; use proto::client::Client; use proto::prepare::PrepareFuture; use proto::query::QueryStream; use proto::typeinfo_composite::TypeinfoCompositeFuture; use proto::typeinfo_enum::TypeinfoEnumFuture; use types::{Kind, Oid, Type}; -use {bad_response, next_statement}; const TYPEINFO_QUERY: &'static str = " SELECT t.typname, t.typtype, t.typelem, r.rngsubtype, t.typbasetype, n.nspname, t.typrelid @@ -46,16 +46,14 @@ pub enum Typeinfo { oid: Oid, client: Client, }, - #[state_machine_future( - transitions( - CachingType, - QueryingEnumVariants, - QueryingDomainBasetype, - QueryingArrayElem, - QueryingCompositeFields, - QueryingRangeSubtype - ) - )] + #[state_machine_future(transitions( + CachingType, + QueryingEnumVariants, + QueryingDomainBasetype, + QueryingArrayElem, + QueryingCompositeFields, + QueryingRangeSubtype + ))] QueryingTypeinfo { future: stream::Collect, oid: Oid, @@ -185,16 +183,30 @@ impl PollTypeinfo for Typeinfo { let row = match rows.get(0) { Some(row) => row, - None => return Err(bad_response()), + None => return Err(Error::unexpected_message()), }; - let name = row.try_get::<_, String>(0)?.ok_or_else(bad_response)?; - let type_ = row.try_get::<_, i8>(1)?.ok_or_else(bad_response)?; - let elem_oid = row.try_get::<_, Oid>(2)?.ok_or_else(bad_response)?; - let rngsubtype = row.try_get::<_, Option>(3)?.ok_or_else(bad_response)?; - let basetype = row.try_get::<_, Oid>(4)?.ok_or_else(bad_response)?; - let schema = row.try_get::<_, String>(5)?.ok_or_else(bad_response)?; - let relid = row.try_get::<_, Oid>(6)?.ok_or_else(bad_response)?; + let name = row + .try_get::<_, String>(0)? + .ok_or_else(Error::unexpected_message)?; + let type_ = row + .try_get::<_, i8>(1)? + .ok_or_else(Error::unexpected_message)?; + let elem_oid = row + .try_get::<_, Oid>(2)? + .ok_or_else(Error::unexpected_message)?; + let rngsubtype = row + .try_get::<_, Option>(3)? + .ok_or_else(Error::unexpected_message)?; + let basetype = row + .try_get::<_, Oid>(4)? + .ok_or_else(Error::unexpected_message)?; + let schema = row + .try_get::<_, String>(5)? + .ok_or_else(Error::unexpected_message)?; + let relid = row + .try_get::<_, Oid>(6)? + .ok_or_else(Error::unexpected_message)?; let kind = if type_ == b'e' as i8 { transition!(QueryingEnumVariants { diff --git a/tokio-postgres/src/proto/typeinfo_composite.rs b/tokio-postgres/src/proto/typeinfo_composite.rs index 127812d4..78257b26 100644 --- a/tokio-postgres/src/proto/typeinfo_composite.rs +++ b/tokio-postgres/src/proto/typeinfo_composite.rs @@ -5,12 +5,12 @@ use std::mem; use std::vec; use error::Error; +use next_statement; use proto::client::Client; use proto::prepare::PrepareFuture; use proto::query::QueryStream; use proto::typeinfo::TypeinfoFuture; use types::{Field, Oid}; -use {bad_response, next_statement}; const TYPEINFO_COMPOSITE_QUERY: &'static str = " SELECT attname, atttypid @@ -95,11 +95,10 @@ impl PollTypeinfoComposite for TypeinfoComposite { let fields = rows .iter() .map(|row| { - let name = row.try_get(0)?.ok_or_else(bad_response)?; - let oid = row.try_get(1)?.ok_or_else(bad_response)?; + let name = row.try_get(0)?.ok_or_else(Error::unexpected_message)?; + let oid = row.try_get(1)?.ok_or_else(Error::unexpected_message)?; Ok((name, oid)) - }) - .collect::, Error>>()?; + }).collect::, Error>>()?; let mut remaining_fields = fields.into_iter(); match remaining_fields.next() { diff --git a/tokio-postgres/src/proto/typeinfo_enum.rs b/tokio-postgres/src/proto/typeinfo_enum.rs index 2f50547a..1058aef9 100644 --- a/tokio-postgres/src/proto/typeinfo_enum.rs +++ b/tokio-postgres/src/proto/typeinfo_enum.rs @@ -3,11 +3,11 @@ use futures::{Async, Future, Poll}; use state_machine_future::RentToOwn; use error::{Error, SqlState}; +use next_statement; use proto::client::Client; use proto::prepare::PrepareFuture; use proto::query::QueryStream; use types::Oid; -use {bad_response, next_statement}; const TYPEINFO_ENUM_QUERY: &'static str = " SELECT enumlabel @@ -126,7 +126,7 @@ impl PollTypeinfoEnum for TypeinfoEnum { let variants = rows .iter() - .map(|row| row.try_get(0)?.ok_or_else(bad_response)) + .map(|row| row.try_get(0)?.ok_or_else(Error::unexpected_message)) .collect::, _>>()?; transition!(Finished((variants, state.client))) diff --git a/tokio-postgres/tests/test.rs b/tokio-postgres/tests/test.rs index f198c6eb..9baa70c7 100644 --- a/tokio-postgres/tests/test.rs +++ b/tokio-postgres/tests/test.rs @@ -50,11 +50,7 @@ fn plain_password_missing() { "postgres://pass_user@localhost:5433".parse().unwrap(), TlsMode::None, ); - match runtime.block_on(handshake) { - Ok(_) => panic!("unexpected success"), - Err(ref e) if e.as_connection().is_some() => {} - Err(e) => panic!("{}", e), - } + runtime.block_on(handshake).err().unwrap(); } #[test] @@ -87,11 +83,7 @@ fn md5_password_missing() { "postgres://md5_user@localhost:5433".parse().unwrap(), TlsMode::None, ); - match runtime.block_on(handshake) { - Ok(_) => panic!("unexpected success"), - Err(ref e) if e.as_connection().is_some() => {} - Err(e) => panic!("{}", e), - } + runtime.block_on(handshake).err().unwrap(); } #[test] @@ -124,11 +116,7 @@ fn scram_password_missing() { "postgres://scram_user@localhost:5433".parse().unwrap(), TlsMode::None, ); - match runtime.block_on(handshake) { - Ok(_) => panic!("unexpected success"), - Err(ref e) if e.as_connection().is_some() => {} - Err(e) => panic!("{}", e), - } + runtime.block_on(handshake).err().unwrap(); } #[test]