use super::builder::StreamingPull;
use super::handler::{AckResult, AtLeastOnce, Handler};
use super::lease_loop::LeaseLoop;
use super::lease_state::LeaseOptions;
use super::leaser::DefaultLeaser;
use super::retry_policy::StreamRetryPolicy;
use super::stream::Stream;
use super::stub::TonicStreaming as _;
use super::transport::Transport;
use crate::google::pubsub::v1::StreamingPullRequest;
use crate::model::Message;
use crate::{Error, Result};
use gaxi::grpc::from_status::to_gax_error;
use gaxi::prost::FromProto as _;
use google_cloud_gax::retry_result::RetryResult;
use std::collections::VecDeque;
use std::sync::Arc;
use tokio::sync::mpsc::UnboundedSender;
#[derive(Debug)]
pub struct Session {
inner: Arc<Transport>,
initial_req: StreamingPullRequest,
stream: Option<Stream<Transport>>,
pool: VecDeque<(Message, Handler)>,
message_tx: UnboundedSender<String>,
ack_tx: UnboundedSender<AckResult>,
_lease_loop: tokio::task::JoinHandle<()>,
}
impl Session {
pub(super) fn new(builder: StreamingPull) -> Self {
let inner = builder.inner;
let subscription = builder.subscription;
let leaser = DefaultLeaser::new(
inner.clone(),
subscription.clone(),
builder.ack_deadline_seconds,
builder.grpc_subchannel_count,
);
let LeaseLoop {
handle: _lease_loop,
message_tx,
ack_tx,
} = LeaseLoop::new(leaser, LeaseOptions::default());
let initial_req = StreamingPullRequest {
subscription,
stream_ack_deadline_seconds: builder.ack_deadline_seconds,
max_outstanding_messages: builder.max_outstanding_messages,
max_outstanding_bytes: builder.max_outstanding_bytes,
client_id: builder.client_id,
protocol_version: 1,
..Default::default()
};
Self {
inner,
initial_req,
stream: None,
pool: VecDeque::new(),
message_tx,
ack_tx,
_lease_loop,
}
}
pub async fn next(&mut self) -> Option<Result<(Message, Handler)>> {
loop {
if let Some(item) = self.pool.pop_front() {
return Some(Ok(item));
}
if let Err(e) = self.read_from_stream().await? {
match StreamRetryPolicy::on_midstream_error(e) {
RetryResult::Continue(_) => {
self.stream = None;
continue;
}
RetryResult::Permanent(e) | RetryResult::Exhausted(e) => {
return Some(Err(e));
}
}
}
}
}
#[cfg(feature = "unstable-stream")]
#[cfg_attr(docsrs, doc(cfg(feature = "unstable-stream")))]
pub fn into_stream(self) -> impl futures::Stream<Item = Result<(Message, Handler)>> + Unpin {
use futures::stream::unfold;
Box::pin(unfold(Some(self), move |state| async move {
if let Some(mut this) = state {
if let Some(chunk) = this.next().await {
return Some((chunk, Some(this)));
}
};
None
}))
}
async fn mut_stream(&mut self) -> Result<&mut Stream<Transport>> {
if self.stream.is_none() {
let stream =
Stream::<Transport>::new(self.inner.clone(), self.initial_req.clone()).await?;
self.stream = Some(stream);
}
Ok(self
.stream
.as_mut()
.expect("`self.stream.is_some()` must be true"))
}
async fn read_from_stream(&mut self) -> Option<Result<()>> {
let resp = {
let stream = match self.mut_stream().await {
Ok(s) => s,
Err(e) => return Some(Err(e)),
};
match stream.next_message().await.transpose()? {
Ok(resp) => resp,
Err(e) => return Some(Err(to_gax_error(e))),
}
};
for rm in resp.received_messages {
let Some(message) = rm.message else {
continue;
};
let _ = self.message_tx.send(rm.ack_id.clone());
let message = match message.cnv().map_err(Error::deser) {
Ok(message) => message,
Err(e) => return Some(Err(e)),
};
self.pool.push_back((
message,
Handler::AtLeastOnce(AtLeastOnce::new(rm.ack_id, self.ack_tx.clone())),
));
}
Some(Ok(()))
}
#[cfg(test)]
async fn close(self) -> anyhow::Result<()> {
drop(self.stream);
drop(self.message_tx);
self._lease_loop.await?;
Ok(())
}
}
#[cfg(test)]
mod tests {
use super::super::client::Subscriber;
use super::super::keepalive::KEEPALIVE_PERIOD;
use super::super::lease_state::tests::{test_id, test_ids};
use super::super::stream::{INITIAL_DELAY, MAXIMUM_DELAY};
use super::*;
use gaxi::grpc::tonic::{Response as TonicResponse, Status as TonicStatus};
use google_cloud_auth::credentials::anonymous::Builder as Anonymous;
use pubsub_grpc_mock::google::pubsub::v1;
use pubsub_grpc_mock::{MockSubscriber, start};
use tokio::sync::mpsc::{channel, unbounded_channel};
use tokio::task::JoinHandle;
use tokio::time::{Duration, Instant};
fn sorted(mut v: Vec<String>) -> Vec<String> {
v.sort();
v
}
fn test_data(v: i32) -> bytes::Bytes {
bytes::Bytes::from(format!("data-{}", test_id(v)))
}
fn test_response(range: std::ops::Range<i32>) -> v1::StreamingPullResponse {
v1::StreamingPullResponse {
received_messages: range
.into_iter()
.map(|i| v1::ReceivedMessage {
ack_id: test_id(i),
message: Some(v1::PubsubMessage {
data: test_data(i).to_vec(),
..Default::default()
}),
..Default::default()
})
.collect(),
..Default::default()
}
}
async fn test_client(endpoint: String) -> anyhow::Result<Subscriber> {
Ok(Subscriber::builder()
.with_endpoint(endpoint)
.with_credentials(Anonymous::new().build())
.build()
.await?)
}
#[tokio::test]
async fn error_starting_stream() -> anyhow::Result<()> {
let mut mock = MockSubscriber::new();
mock.expect_streaming_pull()
.return_once(|_| Err(TonicStatus::failed_precondition("fail")));
let (endpoint, _server) = start("0.0.0.0:0", mock).await?;
let client = test_client(endpoint).await?;
let mut session = client.streaming_pull("projects/p/subscriptions/s").start();
let err = session
.next()
.await
.expect("stream should not be empty")
.expect_err("the first streamed item should be an error");
assert!(err.status().is_some(), "{err:?}");
let status = err.status().unwrap();
assert_eq!(
status.code,
google_cloud_gax::error::rpc::Code::FailedPrecondition
);
assert_eq!(status.message, "fail");
Ok(())
}
#[tokio::test]
async fn initial_request() -> anyhow::Result<()> {
const MIB: i64 = 1024 * 1024;
let (recover_writes_tx, mut recover_writes_rx) = channel(1);
let mut mock = MockSubscriber::new();
mock.expect_streaming_pull().return_once(move |request| {
tokio::spawn(async move {
let mut request_rx = request.into_inner();
while let Some(request) = request_rx.recv().await {
recover_writes_tx
.send(request)
.await
.expect("forwarding writes always succeeds");
}
});
Err(TonicStatus::failed_precondition("fail"))
});
let (endpoint, _server) = start("0.0.0.0:0", mock).await?;
let client = test_client(endpoint).await?;
let _ = client
.streaming_pull("projects/p/subscriptions/s")
.set_ack_deadline_seconds(20)
.set_max_outstanding_messages(2000)
.set_max_outstanding_bytes(200 * MIB)
.start()
.next()
.await;
let initial_req = recover_writes_rx
.recv()
.await
.expect("should receive a request")?;
assert_eq!(initial_req.subscription, "projects/p/subscriptions/s");
assert_eq!(initial_req.stream_ack_deadline_seconds, 20);
assert_eq!(initial_req.max_outstanding_messages, 2000);
assert_eq!(initial_req.max_outstanding_bytes, 200 * MIB);
assert!(
!initial_req.client_id.is_empty(),
"initial request has empty client id: {initial_req:?}"
);
assert!(
initial_req.protocol_version >= 1,
"protocol_version={}",
initial_req.protocol_version
);
Ok(())
}
#[tokio::test(start_paused = true)]
async fn basic_success() -> anyhow::Result<()> {
let (response_tx, response_rx) = channel(10);
let (ack_tx, mut ack_rx) = unbounded_channel();
let mut mock = MockSubscriber::new();
mock.expect_streaming_pull()
.return_once(|_| Ok(TonicResponse::from(response_rx)));
mock.expect_acknowledge().returning(move |r| {
ack_tx
.send(r.into_inner())
.expect("sending on channel always succeeds");
Ok(TonicResponse::from(()))
});
let (endpoint, _server) = start("0.0.0.0:0", mock).await?;
let client = test_client(endpoint).await?;
let mut session = client.streaming_pull("projects/p/subscriptions/s").start();
response_tx.send(Ok(test_response(1..2))).await?;
response_tx.send(Ok(test_response(2..4))).await?;
response_tx.send(Ok(test_response(4..7))).await?;
drop(response_tx);
for i in 1..7 {
let (m, Handler::AtLeastOnce(h)) =
session.next().await.transpose()?.expect("message {i}/6");
assert_eq!(m.data, test_data(i));
assert_eq!(h.ack_id(), test_id(i));
h.ack();
}
let end = session.next().await.transpose()?;
assert!(end.is_none(), "Received extra message: {end:?}");
session.close().await?;
let ack_req = ack_rx.try_recv()?;
assert_eq!(ack_req.subscription, "projects/p/subscriptions/s");
assert_eq!(sorted(ack_req.ack_ids), test_ids(1..7));
Ok(())
}
#[tokio::test(start_paused = true)]
async fn basic_lease_management() -> anyhow::Result<()> {
let (response_tx, response_rx) = channel(10);
let (ack_tx, mut ack_rx) = unbounded_channel();
let (nack_tx, mut nack_rx) = unbounded_channel();
let (extend_tx, mut extend_rx) = unbounded_channel();
let mut mock = MockSubscriber::new();
mock.expect_streaming_pull()
.return_once(|_| Ok(TonicResponse::from(response_rx)));
mock.expect_acknowledge().returning(move |r| {
ack_tx
.send(r.into_inner())
.expect("sending on channel always succeeds");
Ok(TonicResponse::from(()))
});
mock.expect_modify_ack_deadline().returning(move |r| {
let r = r.into_inner();
if r.ack_deadline_seconds == 0 {
nack_tx.send(r).expect("sending on channel always succeeds");
} else {
extend_tx
.send(r)
.expect("sending on channel always succeeds");
}
Ok(TonicResponse::from(()))
});
let (endpoint, _server) = start("0.0.0.0:0", mock).await?;
let client = test_client(endpoint).await?;
let mut session = client.streaming_pull("projects/p/subscriptions/s").start();
response_tx.send(Ok(test_response(0..30))).await?;
drop(response_tx);
for i in 0..10 {
let Some((_, Handler::AtLeastOnce(h))) = session.next().await.transpose()? else {
anyhow::bail!("expected message {i}")
};
h.ack();
}
for i in 10..20 {
let Some((_, Handler::AtLeastOnce(h))) = session.next().await.transpose()? else {
anyhow::bail!("expected message {i}")
};
drop(h);
}
let mut hold = Vec::new();
for i in 20..30 {
let Some((_, Handler::AtLeastOnce(h))) = session.next().await.transpose()? else {
anyhow::bail!("expected message {i}")
};
hold.push(h);
}
tokio::time::advance(Duration::from_secs(10)).await;
session.close().await?;
let ack_req = ack_rx.try_recv()?;
assert_eq!(ack_req.subscription, "projects/p/subscriptions/s");
assert_eq!(sorted(ack_req.ack_ids), test_ids(0..10));
assert!(ack_rx.is_empty(), "{ack_rx:?}");
let nack_req = nack_rx.try_recv()?;
assert_eq!(nack_req.subscription, "projects/p/subscriptions/s");
assert_eq!(nack_req.ack_deadline_seconds, 0);
assert_eq!(sorted(nack_req.ack_ids), test_ids(10..20));
let nack_req = nack_rx.try_recv()?;
assert_eq!(nack_req.subscription, "projects/p/subscriptions/s");
assert_eq!(nack_req.ack_deadline_seconds, 0);
assert_eq!(sorted(nack_req.ack_ids), test_ids(20..30));
assert!(nack_rx.is_empty(), "{nack_rx:?}");
let extend_req = extend_rx.try_recv()?;
assert_eq!(extend_req.subscription, "projects/p/subscriptions/s");
assert_eq!(extend_req.ack_deadline_seconds, 10);
assert_eq!(sorted(extend_req.ack_ids), test_ids(20..30));
Ok(())
}
#[tokio::test(start_paused = true)]
async fn delayed_responses() -> anyhow::Result<()> {
let (response_tx, response_rx) = channel(10);
let handle: JoinHandle<anyhow::Result<()>> = tokio::spawn(async move {
tokio::time::sleep(Duration::from_millis(20)).await;
response_tx.send(Ok(test_response(1..2))).await?;
Ok(())
});
let mut mock = MockSubscriber::new();
mock.expect_streaming_pull()
.return_once(|_| Ok(TonicResponse::from(response_rx)));
let (endpoint, _server) = start("0.0.0.0:0", mock).await?;
let client = test_client(endpoint).await?;
let mut session = client.streaming_pull("projects/p/subscriptions/s").start();
let (m, Handler::AtLeastOnce(h)) = session
.next()
.await
.transpose()?
.expect("stream should wait for a message");
assert_eq!(m.data, test_data(1));
assert_eq!(h.ack_id(), test_id(1));
handle.await??;
Ok(())
}
#[tokio::test]
async fn serves_messages_immediately() -> anyhow::Result<()> {
let (response_tx, response_rx) = channel(10);
let mut mock = MockSubscriber::new();
mock.expect_streaming_pull()
.return_once(|_| Ok(TonicResponse::from(response_rx)));
let (endpoint, _server) = start("0.0.0.0:0", mock).await?;
let client = test_client(endpoint).await?;
let mut session = client.streaming_pull("projects/p/subscriptions/s").start();
for i in 1..7 {
response_tx.send(Ok(test_response(i..i + 1))).await?;
let (m, Handler::AtLeastOnce(h)) =
session.next().await.transpose()?.expect("message {i}/6");
assert_eq!(m.data, test_data(i));
assert_eq!(h.ack_id(), test_id(i));
}
drop(response_tx);
let end = session.next().await.transpose()?;
assert!(end.is_none(), "Received extra message: {end:?}");
Ok(())
}
#[tokio::test]
async fn handles_empty_response() -> anyhow::Result<()> {
let (response_tx, response_rx) = channel(10);
let mut mock = MockSubscriber::new();
mock.expect_streaming_pull()
.return_once(|_| Ok(TonicResponse::from(response_rx)));
let (endpoint, _server) = start("0.0.0.0:0", mock).await?;
let client = test_client(endpoint).await?;
let mut session = client.streaming_pull("projects/p/subscriptions/s").start();
response_tx.send(Ok(test_response(1..2))).await?;
response_tx.send(Ok(test_response(2..2))).await?;
response_tx.send(Ok(test_response(2..3))).await?;
drop(response_tx);
for i in 1..3 {
let (m, Handler::AtLeastOnce(h)) =
session.next().await.transpose()?.expect("message {i}/2");
assert_eq!(m.data, test_data(i));
assert_eq!(h.ack_id(), test_id(i));
}
let end = session.next().await.transpose()?;
assert!(end.is_none(), "Received extra message: {end:?}");
Ok(())
}
#[tokio::test(start_paused = true)]
async fn handles_missing_message_field() -> anyhow::Result<()> {
let (response_tx, response_rx) = channel(10);
let (extend_tx, mut extend_rx) = unbounded_channel();
let bad = v1::StreamingPullResponse {
received_messages: vec![v1::ReceivedMessage {
ack_id: "ignored-ack-id".to_string(),
message: None,
..Default::default()
}],
..Default::default()
};
let mut mock = MockSubscriber::new();
mock.expect_streaming_pull()
.return_once(|_| Ok(TonicResponse::from(response_rx)));
mock.expect_acknowledge()
.returning(|_| Ok(TonicResponse::from(())));
mock.expect_modify_ack_deadline().returning(move |r| {
extend_tx
.send(r.into_inner())
.expect("sending on channel always succeeds");
Ok(TonicResponse::from(()))
});
let (endpoint, _server) = start("0.0.0.0:0", mock).await?;
let client = test_client(endpoint).await?;
let mut session = client.streaming_pull("projects/p/subscriptions/s").start();
response_tx.send(Ok(test_response(1..4))).await?;
response_tx.send(Ok(bad)).await?;
response_tx.send(Ok(test_response(4..7))).await?;
drop(response_tx);
let mut handlers = Vec::new();
for i in 1..7 {
let (m, Handler::AtLeastOnce(h)) =
session.next().await.transpose()?.expect("message {i}/6");
assert_eq!(m.data, test_data(i));
assert_eq!(h.ack_id(), test_id(i));
handlers.push(h);
}
let end = session.next().await.transpose()?;
assert!(end.is_none(), "Received extra message: {end:?}");
tokio::time::advance(Duration::from_secs(10)).await;
session.close().await?;
let extend_req = extend_rx.try_recv()?;
assert_eq!(extend_req.subscription, "projects/p/subscriptions/s");
assert_eq!(extend_req.ack_deadline_seconds, 10);
assert_eq!(sorted(extend_req.ack_ids), test_ids(1..7));
Ok(())
}
#[tokio::test]
async fn permanent_error_midstream() -> anyhow::Result<()> {
let (response_tx, response_rx) = channel(10);
let mut mock = MockSubscriber::new();
mock.expect_streaming_pull()
.return_once(|_| Ok(TonicResponse::from(response_rx)));
let (endpoint, _server) = start("0.0.0.0:0", mock).await?;
let client = test_client(endpoint).await?;
let mut session = client.streaming_pull("projects/p/subscriptions/s").start();
response_tx.send(Ok(test_response(1..4))).await?;
response_tx
.send(Err(TonicStatus::failed_precondition("fail")))
.await?;
drop(response_tx);
for i in 1..4 {
let (m, Handler::AtLeastOnce(h)) =
session.next().await.transpose()?.expect("message {i}/3");
assert_eq!(m.data, test_data(i));
assert_eq!(h.ack_id(), test_id(i));
}
let err = session
.next()
.await
.transpose()
.expect_err("expected an error from stream");
assert!(err.status().is_some(), "{err:?}");
let status = err.status().unwrap();
assert_eq!(
status.code,
google_cloud_gax::error::rpc::Code::FailedPrecondition
);
assert_eq!(status.message, "fail");
Ok(())
}
#[tokio::test(start_paused = true)]
async fn keepalives() -> anyhow::Result<()> {
let (recover_writes_tx, mut recover_writes_rx) = channel(1);
let (response_tx, response_rx) = channel(10);
let mut mock = MockSubscriber::new();
mock.expect_streaming_pull().return_once(move |request| {
tokio::spawn(async move {
let mut request_rx = request.into_inner();
while let Some(request) = request_rx.recv().await {
recover_writes_tx
.send(request)
.await
.expect("forwarding writes always succeeds");
}
});
Ok(TonicResponse::from(response_rx))
});
mock.expect_modify_ack_deadline()
.returning(|_| Ok(TonicResponse::from(())));
let (endpoint, _server) = start("0.0.0.0:0", mock).await?;
let client = test_client(endpoint).await?;
let mut session = client.streaming_pull("projects/p/subscriptions/s").start();
response_tx.send(Ok(test_response(1..4))).await?;
let _ = session.next().await;
let initial_req = recover_writes_rx
.recv()
.await
.expect("should receive an initial request")?;
assert_eq!(initial_req.subscription, "projects/p/subscriptions/s");
tokio::time::advance(KEEPALIVE_PERIOD).await;
let keepalive_req = recover_writes_rx
.recv()
.await
.expect("should receive a keepalive request")?;
assert_eq!(keepalive_req, v1::StreamingPullRequest::default());
drop(session);
tokio::time::advance(4 * KEEPALIVE_PERIOD).await;
assert!(recover_writes_rx.is_empty(), "{recover_writes_rx:?}");
Ok(())
}
#[tokio::test]
async fn client_id() -> anyhow::Result<()> {
let (recover_writes_tx, mut recover_writes_rx) = channel(10);
let recover_writes_tx = std::sync::Arc::new(tokio::sync::Mutex::new(recover_writes_tx));
let mut mock = MockSubscriber::new();
mock.expect_streaming_pull()
.times(3)
.returning(move |request| {
let tx = recover_writes_tx.clone();
tokio::spawn(async move {
let mut request_rx = request.into_inner();
while let Some(request) = request_rx.recv().await {
tx.lock()
.await
.send(request)
.await
.expect("forwarding writes always succeeds");
}
});
Err(TonicStatus::failed_precondition("fail"))
});
let (endpoint, _server) = start("0.0.0.0:0", mock).await?;
let c1 = test_client(endpoint.clone()).await?;
let _ = c1
.streaming_pull("projects/p/subscriptions/s")
.start()
.next()
.await;
let req1 = recover_writes_rx
.recv()
.await
.expect("should receive a request")?;
let _ = c1
.streaming_pull("projects/p/subscriptions/s")
.start()
.next()
.await;
let req2 = recover_writes_rx
.recv()
.await
.expect("should receive a request")?;
assert_eq!(req1.client_id, req2.client_id);
let c2 = test_client(endpoint).await?;
let _ = c2
.streaming_pull("projects/p/subscriptions/s")
.start()
.next()
.await;
let req3 = recover_writes_rx
.recv()
.await
.expect("should receive a request")?;
assert_ne!(req1.client_id, req3.client_id);
Ok(())
}
#[tokio::test(start_paused = true)]
async fn no_immediate_message() -> anyhow::Result<()> {
const TEST_TIMEOUT: Duration = Duration::from_secs(42);
let (_response_tx, response_rx) = channel(10);
let mut mock = MockSubscriber::new();
mock.expect_streaming_pull()
.return_once(move |_| Ok(TonicResponse::from(response_rx)));
let (endpoint, _server) = start("0.0.0.0:0", mock).await?;
let client = test_client(endpoint).await?;
let mut session = client.streaming_pull("projects/p/subscriptions/s").start();
let _ = tokio::time::timeout(TEST_TIMEOUT, session.next())
.await
.expect_err("next() should never yield.");
Ok(())
}
#[tokio::test(start_paused = true)]
async fn retry_transient_when_starting_stream() -> anyhow::Result<()> {
const NUM_RETRIES: u32 = 20;
let start_time = Instant::now();
let mut seq = mockall::Sequence::new();
let mut mock = MockSubscriber::new();
mock.expect_streaming_pull()
.times(NUM_RETRIES as usize)
.in_sequence(&mut seq)
.returning(|_| Err(TonicStatus::unavailable("try again")));
mock.expect_streaming_pull()
.times(1)
.in_sequence(&mut seq)
.return_once(|_| Err(TonicStatus::failed_precondition("fail")));
let (endpoint, _server) = start("0.0.0.0:0", mock).await?;
let client = test_client(endpoint).await?;
let mut session = client.streaming_pull("projects/p/subscriptions/s").start();
let err = session
.next()
.await
.expect("stream should not be empty")
.expect_err("the first streamed item should be an error");
assert!(err.status().is_some(), "{err:?}");
let status = err.status().unwrap();
assert_eq!(
status.code,
google_cloud_gax::error::rpc::Code::FailedPrecondition
);
assert_eq!(status.message, "fail");
let elapsed = start_time.elapsed();
assert!(
elapsed <= MAXIMUM_DELAY * NUM_RETRIES,
"elapsed={elapsed:?}"
);
assert!(
elapsed >= INITIAL_DELAY * NUM_RETRIES,
"elapsed={elapsed:?}"
);
Ok(())
}
#[tokio::test(start_paused = true)]
async fn resume_midstream_success() -> anyhow::Result<()> {
let (response_tx_1, response_rx_1) = channel(10);
let (response_tx_2, response_rx_2) = channel(10);
let (response_tx_3, response_rx_3) = channel(10);
let (ack_tx, mut ack_rx) = unbounded_channel();
let mut seq = mockall::Sequence::new();
let mut mock = MockSubscriber::new();
mock.expect_streaming_pull()
.times(1)
.in_sequence(&mut seq)
.return_once(|_| Ok(TonicResponse::from(response_rx_1)));
mock.expect_streaming_pull()
.times(1)
.in_sequence(&mut seq)
.return_once(move |_| Ok(TonicResponse::from(response_rx_2)));
mock.expect_streaming_pull()
.times(1)
.in_sequence(&mut seq)
.return_once(|_| Ok(TonicResponse::from(response_rx_3)));
mock.expect_acknowledge().times(1..).returning(move |r| {
ack_tx
.send(r.into_inner())
.expect("sending on channel always succeeds");
Ok(TonicResponse::from(()))
});
let (endpoint, _server) = start("0.0.0.0:0", mock).await?;
let client = test_client(endpoint).await?;
let mut session = client.streaming_pull("projects/p/subscriptions/s").start();
response_tx_1.send(Ok(test_response(0..10))).await?;
response_tx_1.send(Ok(test_response(10..20))).await?;
response_tx_1
.send(Err(TonicStatus::unavailable("GFE disconnect. try again")))
.await?;
drop(response_tx_1);
response_tx_2.send(Ok(test_response(20..30))).await?;
response_tx_2.send(Ok(test_response(30..40))).await?;
response_tx_2
.send(Err(TonicStatus::unavailable("GFE disconnect. try again")))
.await?;
drop(response_tx_2);
response_tx_3.send(Ok(test_response(40..50))).await?;
drop(response_tx_3);
for i in 0..50 {
let (m, h) = session
.next()
.await
.unwrap_or_else(|| panic!("expected message {}/50", i + 1))?;
assert_eq!(m.data, test_data(i));
h.ack();
}
let end = session.next().await.transpose()?;
assert!(end.is_none(), "Received extra message: {end:?}");
session.close().await?;
let mut got = Vec::new();
while let Ok(ack_req) = ack_rx.try_recv() {
assert_eq!(ack_req.subscription, "projects/p/subscriptions/s");
got.extend(ack_req.ack_ids);
}
assert_eq!(sorted(got), test_ids(0..50));
Ok(())
}
#[tokio::test(start_paused = true)]
async fn resume_midstream_hits_permanent_error() -> anyhow::Result<()> {
let (response_tx, response_rx) = channel(10);
let (ack_tx, mut ack_rx) = unbounded_channel();
let mut seq = mockall::Sequence::new();
let mut mock = MockSubscriber::new();
mock.expect_streaming_pull()
.times(1)
.in_sequence(&mut seq)
.return_once(|_| Ok(TonicResponse::from(response_rx)));
mock.expect_streaming_pull()
.times(3)
.in_sequence(&mut seq)
.returning(|_| Err(TonicStatus::unavailable("try again")));
mock.expect_streaming_pull()
.times(1)
.in_sequence(&mut seq)
.return_once(|_| Err(TonicStatus::failed_precondition("fail")));
mock.expect_acknowledge().times(1..).returning(move |r| {
ack_tx
.send(r.into_inner())
.expect("sending on channel always succeeds");
Ok(TonicResponse::from(()))
});
let (endpoint, _server) = start("0.0.0.0:0", mock).await?;
let client = test_client(endpoint).await?;
let mut session = client.streaming_pull("projects/p/subscriptions/s").start();
response_tx.send(Ok(test_response(0..10))).await?;
response_tx.send(Ok(test_response(10..20))).await?;
response_tx
.send(Err(TonicStatus::unavailable("GFE disconnect. try again")))
.await?;
drop(response_tx);
for i in 0..20 {
let (m, h) = session
.next()
.await
.unwrap_or_else(|| panic!("expected message {}/20", i + 1))?;
assert_eq!(m.data, test_data(i));
h.ack();
}
let err = session
.next()
.await
.transpose()
.expect_err("expected an error from stream");
assert!(err.status().is_some(), "{err:?}");
let status = err.status().unwrap();
assert_eq!(
status.code,
google_cloud_gax::error::rpc::Code::FailedPrecondition
);
assert_eq!(status.message, "fail");
session.close().await?;
let mut got = Vec::new();
while let Ok(ack_req) = ack_rx.try_recv() {
assert_eq!(ack_req.subscription, "projects/p/subscriptions/s");
got.extend(ack_req.ack_ids);
}
assert_eq!(sorted(got), test_ids(0..20));
Ok(())
}
#[tokio::test]
async fn routing_header() -> anyhow::Result<()> {
let mut mock = MockSubscriber::new();
mock.expect_streaming_pull().return_once(move |request| {
let metadata = request.metadata();
assert_eq!(
metadata
.get("x-goog-request-params")
.expect("routing header missing"),
"subscription=projects/p/subscriptions/s"
);
Err(TonicStatus::failed_precondition("ignored"))
});
let (endpoint, _server) = start("0.0.0.0:0", mock).await?;
let client = test_client(endpoint).await?;
let _ = client
.streaming_pull("projects/p/subscriptions/s")
.start()
.next()
.await;
Ok(())
}
#[cfg(feature = "unstable-stream")]
#[tokio::test(start_paused = true)]
async fn into_stream() -> anyhow::Result<()> {
use futures::TryStreamExt;
let (response_tx, response_rx) = channel(10);
let (ack_tx, mut ack_rx) = unbounded_channel();
let mut mock = MockSubscriber::new();
mock.expect_streaming_pull()
.return_once(|_| Ok(TonicResponse::from(response_rx)));
mock.expect_acknowledge().returning(move |r| {
ack_tx
.send(r.into_inner())
.expect("sending on channel always succeeds");
Ok(TonicResponse::from(()))
});
let (endpoint, _server) = start("0.0.0.0:0", mock).await?;
let client = test_client(endpoint).await?;
let stream = client
.streaming_pull("projects/p/subscriptions/s")
.start()
.into_stream();
response_tx.send(Ok(test_response(1..3))).await?;
drop(response_tx);
let got: Vec<_> = stream
.map_ok(|(m, h)| {
h.ack();
m.data
})
.try_collect()
.await?;
assert_eq!(got, vec![test_data(1), test_data(2)]);
let ack_req = ack_rx
.recv()
.await
.expect("should receive acknowledgements");
assert_eq!(ack_req.subscription, "projects/p/subscriptions/s");
assert_eq!(sorted(ack_req.ack_ids), test_ids(1..3));
Ok(())
}
}