From 9126ec4ef21e0903bbcaed36809c6ebf55c766da Mon Sep 17 00:00:00 2001 From: Steven Fackler Date: Thu, 2 Mar 2017 22:27:12 -0800 Subject: [PATCH] Support notifications on tokio-postgres Closes #242 --- postgres-shared/src/lib.rs | 11 +++++ postgres/src/notification.rs | 14 ++---- tokio-postgres/src/lib.rs | 90 +++++++++++++++++++++++++++++------- tokio-postgres/src/test.rs | 56 +++++++++++++++------- 4 files changed, 128 insertions(+), 43 deletions(-) diff --git a/postgres-shared/src/lib.rs b/postgres-shared/src/lib.rs index 8218efb5..521779b3 100644 --- a/postgres-shared/src/lib.rs +++ b/postgres-shared/src/lib.rs @@ -19,3 +19,14 @@ pub struct CancelData { /// The secret key for the session. pub secret_key: i32, } + +/// An asynchronous notification. +#[derive(Clone, Debug)] +pub struct Notification { + /// The process ID of the notifying backend process. + pub process_id: i32, + /// The name of the channel that the notify has been raised on. + pub channel: String, + /// The "payload" string passed from the notifying process. + pub payload: String, +} diff --git a/postgres/src/notification.rs b/postgres/src/notification.rs index 8fedede0..442b331c 100644 --- a/postgres/src/notification.rs +++ b/postgres/src/notification.rs @@ -5,20 +5,12 @@ use std::fmt; use std::time::Duration; use postgres_protocol::message::backend; +#[doc(inline)] +pub use postgres_shared::Notification; + use {desynchronized, Result, Connection, NotificationsNew}; use error::Error; -/// An asynchronous notification. -#[derive(Clone, Debug)] -pub struct Notification { - /// The process ID of the notifying backend process. - pub process_id: i32, - /// The name of the channel that the notify has been raised on. - pub channel: String, - /// The "payload" string passed from the notifying process. - pub payload: String, -} - /// Notifications from the Postgres backend. pub struct Notifications<'conn> { conn: &'conn Connection, diff --git a/tokio-postgres/src/lib.rs b/tokio-postgres/src/lib.rs index 9e89b37f..5f827721 100644 --- a/tokio-postgres/src/lib.rs +++ b/tokio-postgres/src/lib.rs @@ -55,13 +55,15 @@ #![warn(missing_docs)] extern crate fallible_iterator; -extern crate futures; extern crate futures_state_stream; extern crate postgres_shared; extern crate postgres_protocol; extern crate tokio_core; extern crate tokio_dns; +#[macro_use] +extern crate futures; + #[cfg(unix)] extern crate tokio_uds; @@ -71,23 +73,24 @@ extern crate tokio_openssl; extern crate openssl; use fallible_iterator::FallibleIterator; -use futures::{Future, IntoFuture, BoxFuture, Stream, Sink, Poll, StartSend}; +use futures::{Future, IntoFuture, BoxFuture, Stream, Sink, Poll, StartSend, Async}; use futures::future::Either; use futures_state_stream::{StreamEvent, StateStream, BoxStateStream, FutureExt}; use postgres_protocol::authentication; use postgres_protocol::message::{backend, frontend}; use postgres_protocol::message::backend::{ErrorResponseBody, ErrorFields}; use postgres_shared::rows::RowData; -use std::collections::HashMap; +use std::collections::{HashMap, VecDeque}; use std::fmt; use std::io; use std::sync::Arc; use std::sync::atomic::{AtomicUsize, ATOMIC_USIZE_INIT, Ordering}; use std::sync::mpsc::{self, Sender, Receiver}; +use tokio_core::io::IoFuture; use tokio_core::reactor::Handle; #[doc(inline)] -pub use postgres_shared::{params, CancelData}; +pub use postgres_shared::{params, CancelData, Notification}; use error::{ConnectError, Error, DbError, SqlState}; use params::{ConnectParams, IntoConnectParams}; @@ -166,6 +169,7 @@ struct InnerConnection { close_sender: Sender<(u8, String)>, parameters: HashMap, types: HashMap, + notifications: VecDeque, cancel_data: CancelData, has_typeinfo_query: bool, has_typeinfo_enum_query: bool, @@ -173,25 +177,27 @@ struct InnerConnection { } impl InnerConnection { - fn read(self) -> BoxFuture<(backend::Message>, InnerConnection), io::Error> { + fn read(self) -> IoFuture<(backend::Message>, InnerConnection)> { self.into_future() .map_err(|e| e.0) .and_then(|(m, mut s)| { match m { - Some(backend::Message::ParameterStatus(body)) => { - let name = match body.name() { - Ok(name) => name.to_owned(), + Some(backend::Message::NotificationResponse(body)) => { + let process_id = body.process_id(); + let channel = match body.channel() { + Ok(channel) => channel.to_owned(), Err(e) => return Either::A(Err(e).into_future()), }; - let value = match body.value() { - Ok(value) => value.to_owned(), + let message = match body.message() { + Ok(channel) => channel.to_owned(), Err(e) => return Either::A(Err(e).into_future()), }; - s.parameters.insert(name, value); - Either::B(s.read()) - } - Some(backend::Message::NoticeResponse(_)) => { - // TODO forward the error + let notification = Notification { + process_id: process_id, + channel: channel, + payload: message, + }; + s.notifications.push_back(notification); Either::B(s.read()) } Some(m) => Either::A(Ok((m, s)).into_future()), @@ -210,7 +216,18 @@ impl Stream for InnerConnection { type Error = io::Error; fn poll(&mut self) -> Poll>>, io::Error> { - self.stream.poll() + loop { + match try_ready!(self.stream.poll()) { + Some(backend::Message::ParameterStatus(body)) => { + let name = body.name()?.to_owned(); + let value = body.value()?.to_owned(); + self.parameters.insert(name, value); + } + // TODO forward to a handler + Some(backend::Message::NoticeResponse(_)) => {} + msg => return Ok(Async::Ready(msg)), + } + } } } @@ -279,6 +296,7 @@ impl Connection { close_receiver: receiver, parameters: HashMap::new(), types: HashMap::new(), + notifications: VecDeque::new(), cancel_data: CancelData { process_id: 0, secret_key: 0, @@ -1000,6 +1018,11 @@ impl Connection { .boxed() } + /// Returns a stream of asynchronus notifications receieved from the server. + pub fn notifications(self) -> Notifications { + Notifications(self) + } + /// Returns information used to cancel pending queries. /// /// Used with the `cancel_query` function. The object returned can be used @@ -1015,6 +1038,41 @@ impl Connection { } } +/// A stream of asynchronous Postgres notifications. +pub struct Notifications(Connection); + +impl Notifications { + /// Consumes the `Notifications`, returning the inner `Connection`. + pub fn into_inner(self) -> Connection { + self.0 + } +} + +impl Stream for Notifications { + type Item = Notification; + + type Error = Error; + + fn poll(&mut self) -> Poll, Error> { + if let Some(notification) = (self.0).0.notifications.pop_front() { + return Ok(Async::Ready(Some(notification))); + } + + match try_ready!((self.0).0.poll()) { + Some(backend::Message::NotificationResponse(body)) => { + let notification = Notification { + process_id: body.process_id(), + channel: body.channel()?.to_owned(), + payload: body.message()?.to_owned(), + }; + Ok(Async::Ready(Some(notification))) + } + Some(_) => Err(bad_message()), + None => Ok(Async::Ready(None)), + } + } +} + fn connect_err(fields: &mut ErrorFields) -> ConnectError { match DbError::new(fields) { Ok(err) => ConnectError::Db(Box::new(err)), diff --git a/tokio-postgres/src/test.rs b/tokio-postgres/src/test.rs index 5614c4b9..1d5fe95c 100644 --- a/tokio-postgres/src/test.rs +++ b/tokio-postgres/src/test.rs @@ -157,22 +157,23 @@ fn query() { #[test] fn transaction() { let mut l = Core::new().unwrap(); - let done = Connection::connect("postgres://postgres@localhost", TlsMode::None, &l.handle()) - .then(|c| { - c.unwrap().batch_execute("CREATE TEMPORARY TABLE foo (id SERIAL, name VARCHAR);") - }) - .then(|c| c.unwrap().transaction()) - .then(|t| t.unwrap().batch_execute("INSERT INTO foo (name) VALUES ('joe');")) - .then(|t| t.unwrap().rollback()) - .then(|c| c.unwrap().transaction()) - .then(|t| t.unwrap().batch_execute("INSERT INTO foo (name) VALUES ('bob');")) - .then(|t| t.unwrap().commit()) - .then(|c| c.unwrap().prepare("SELECT name FROM foo")) - .and_then(|(s, c)| c.query(&s, &[]).collect()) - .map(|(r, _)| { - assert_eq!(r.len(), 1); - assert_eq!(r[0].get::("name"), "bob"); - }); + let done = + Connection::connect("postgres://postgres@localhost", TlsMode::None, &l.handle()) + .then(|c| { + c.unwrap().batch_execute("CREATE TEMPORARY TABLE foo (id SERIAL, name VARCHAR);") + }) + .then(|c| c.unwrap().transaction()) + .then(|t| t.unwrap().batch_execute("INSERT INTO foo (name) VALUES ('joe');")) + .then(|t| t.unwrap().rollback()) + .then(|c| c.unwrap().transaction()) + .then(|t| t.unwrap().batch_execute("INSERT INTO foo (name) VALUES ('bob');")) + .then(|t| t.unwrap().commit()) + .then(|c| c.unwrap().prepare("SELECT name FROM foo")) + .and_then(|(s, c)| c.query(&s, &[]).collect()) + .map(|(r, _)| { + assert_eq!(r.len(), 1); + assert_eq!(r[0].get::("name"), "bob"); + }); l.run(done).unwrap(); } @@ -382,3 +383,26 @@ fn cancel() { Ok(_) => panic!("unexpected success"), } } + +#[test] +fn notifications() { + let mut l = Core::new().unwrap(); + let handle = l.handle(); + + let done = Connection::connect("postgres://postgres@localhost", TlsMode::None, &handle) + .then(|c| c.unwrap().batch_execute("LISTEN test_notifications")) + .and_then(|c1| { + Connection::connect("postgres://postgres@localhost", TlsMode::None, &handle) + .then(|c2| { + c2.unwrap().batch_execute("NOTIFY test_notifications, 'foo'").map(|_| c1) + }) + }) + .and_then(|c| c.notifications().into_future().map_err(|(e, _)| e)) + .map(|(n, _)| { + let n = n.unwrap(); + assert_eq!(n.channel, "test_notifications"); + assert_eq!(n.payload, "foo"); + }); + + l.run(done).unwrap(); +}