Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 2 additions & 0 deletions src/gax-internal/src/grpc/grpc_rust.rs
Original file line number Diff line number Diff line change
Expand Up @@ -42,6 +42,8 @@ pub mod bidi;
mod metadata;
mod receive;
mod send;
#[cfg(test)]
pub(crate) mod testing;
mod unary;

pub use bidi::GrpcRustStreaming;
Expand Down
108 changes: 18 additions & 90 deletions src/gax-internal/src/grpc/grpc_rust/receive.rs
Original file line number Diff line number Diff line change
Expand Up @@ -239,6 +239,7 @@ fn grpc_rust_error_to_tonic_code(code: StatusCodeError) -> tonic::Code {

#[cfg(test)]
mod tests {
use super::super::testing::*;
use super::*;
use grpc::StatusError;
use grpc::client::ResponseHeaders;
Expand Down Expand Up @@ -300,18 +301,10 @@ mod tests {
assert_eq!(grpc_rust_error_to_tonic_code(input), want);
}

#[derive(Clone, PartialEq, prost::Message)]
struct TestMessage {
#[prost(string, tag = "1")]
value: String,
}

#[test]
fn grpc_rust_recv_decodes_correctly() -> anyhow::Result<()> {
// Arrange
let want = TestMessage {
value: "hello".to_string(),
};
let want = TestMessage::new("hello");
let mut encoded = bytes::Bytes::from(want.encode_to_vec());
let mut recv = GrpcRustRecv::<TestMessage>::new();

Expand All @@ -324,27 +317,11 @@ mod tests {
Ok(())
}

struct TestClosedStream;

impl RecvStream for TestClosedStream {
async fn recv(&mut self, _buf: &mut dyn RecvMessage) -> ResponseStreamItem {
ResponseStreamItem::StreamClosed
}
}

struct TestPendingStream;

impl RecvStream for TestPendingStream {
async fn recv(&mut self, _buf: &mut dyn RecvMessage) -> ResponseStreamItem {
std::future::pending().await
}
}

#[tokio::test]
async fn receive_task_join_returns_internal_status_when_stream_closed_without_trailers()
-> anyhow::Result<()> {
// Arrange
let (mut rx, mut task) = ReceiveTask::start::<TestMessage, _>(TestClosedStream);
let (mut rx, mut task) = ReceiveTask::start::<TestMessage, _>(ClosedRecvStream);

// Act
let item = rx
Expand All @@ -369,7 +346,7 @@ mod tests {
async fn receive_task_join_returns_internal_status_when_receiver_dropped_early()
-> anyhow::Result<()> {
// Arrange
let (rx, mut task) = ReceiveTask::start::<TestMessage, _>(TestPendingStream);
let (rx, mut task) = ReceiveTask::start::<TestMessage, _>(PendingRecvStream);
drop(rx);

// Act
Expand All @@ -386,15 +363,8 @@ mod tests {
// Arrange
const SIMULATED_PANIC_MSG: &str = "simulated panic in recv stream";

struct TestPanicStream;

impl RecvStream for TestPanicStream {
async fn recv(&mut self, _buf: &mut dyn RecvMessage) -> ResponseStreamItem {
panic!("{SIMULATED_PANIC_MSG}");
}
}

let (_rx, mut task) = ReceiveTask::start::<TestMessage, _>(TestPanicStream);
let (_rx, mut task) =
ReceiveTask::start::<TestMessage, _>(PanicRecvStream::new(SIMULATED_PANIC_MSG));

// Act
let status = task.join().await;
Expand All @@ -408,7 +378,7 @@ mod tests {
#[tokio::test]
async fn receive_task_join_returns_cancelled_status_when_task_aborted() -> anyhow::Result<()> {
// Arrange
let (_rx, mut task) = ReceiveTask::start::<TestMessage, _>(TestPendingStream);
let (_rx, mut task) = ReceiveTask::start::<TestMessage, _>(PendingRecvStream);
let handle = task
.handle
.as_ref()
Expand All @@ -427,7 +397,7 @@ mod tests {
#[tokio::test]
async fn receive_task_join_clears_handle_allowing_safe_drop() -> anyhow::Result<()> {
// Arrange
let (_rx, mut task) = ReceiveTask::start::<TestMessage, _>(TestClosedStream);
let (_rx, mut task) = ReceiveTask::start::<TestMessage, _>(ClosedRecvStream);

// Act
let status = task.join().await;
Expand All @@ -442,7 +412,7 @@ mod tests {
#[tokio::test]
async fn receive_task_join_called_multiple_times_returns_error() -> anyhow::Result<()> {
// Arrange
let (_rx, mut task) = ReceiveTask::start::<TestMessage, _>(TestPendingStream);
let (_rx, mut task) = ReceiveTask::start::<TestMessage, _>(PendingRecvStream);
let handle = task
.handle
.as_ref()
Expand Down Expand Up @@ -483,52 +453,15 @@ mod tests {
const ERROR_MESSAGE_CHANNEL_YIELD: &str = "channel should yield item";
const ERROR_MESSAGE_NON_TERMINAL: &str = "expected non-terminal item";

// TODO(#5991): Refactor common stream state test mocks across grpc_rust tests.
#[derive(Default)]
enum StreamState {
#[default]
Initial,
HeadersSent,
MessageSent,
Done,
}

struct TestHeadersStream {
state: StreamState,
}

impl RecvStream for TestHeadersStream {
async fn recv(&mut self, message: &mut dyn RecvMessage) -> ResponseStreamItem {
match self.state {
StreamState::Initial => {
self.state = StreamState::HeadersSent;
let mut metadata = grpc::metadata::MetadataMap::new();
metadata.insert(HEADER_KEY, MetadataValue::from_static(HEADER_VALUE));
ResponseStreamItem::Headers(ResponseHeaders::new().with_metadata(metadata))
}
StreamState::HeadersSent => {
self.state = StreamState::MessageSent;
let response = TestMessage {
value: MESSAGE_VALUE.to_string(),
};
let mut encoded = bytes::Bytes::from(response.encode_to_vec());
message
.decode(&mut encoded)
.expect("decode response message");
ResponseStreamItem::Message
}
StreamState::MessageSent => {
self.state = StreamState::Done;
ResponseStreamItem::Trailers(Trailers::new(Ok(())))
}
StreamState::Done => ResponseStreamItem::StreamClosed,
}
}
}
let mut metadata = grpc::metadata::MetadataMap::new();
metadata.insert(HEADER_KEY, MetadataValue::from_static(HEADER_VALUE));
let stream = MockRecvStream::new([
MockRecvAction::Headers(ResponseHeaders::new().with_metadata(metadata)),
MockRecvAction::Message(TestMessage::new(MESSAGE_VALUE)),
MockRecvAction::Trailers(Trailers::new(Ok(()))),
]);

let (mut rx, _task) = ReceiveTask::start::<TestMessage, _>(TestHeadersStream {
state: StreamState::default(),
});
let (mut rx, _task) = ReceiveTask::start::<TestMessage, _>(stream);

// Act
let first_item = rx
Expand Down Expand Up @@ -559,12 +492,7 @@ mod tests {
RecvItem::Headers(_) => panic!("expected message item"),
RecvItem::Message(m) => m,
};
assert_eq!(
msg,
TestMessage {
value: MESSAGE_VALUE.to_string()
}
);
assert_eq!(msg, TestMessage::new(MESSAGE_VALUE));

// Act
let terminal_item = rx
Expand Down
89 changes: 18 additions & 71 deletions src/gax-internal/src/grpc/grpc_rust/send.rs
Original file line number Diff line number Diff line change
Expand Up @@ -197,45 +197,14 @@ where

#[cfg(test)]
mod tests {
use super::super::testing::*;
use super::*;
use pretty_assertions::assert_eq;

#[derive(Clone, PartialEq, prost::Message)]
struct TestMessage {
#[prost(string, tag = "1")]
value: String,
}

struct TestPendingSendStream;

impl SendStream for TestPendingSendStream {
async fn send(
&mut self,
_message: &dyn SendMessage,
_options: SendOptions,
) -> Result<(), ()> {
std::future::pending().await
}
}

struct TestFailingSendStream;

impl SendStream for TestFailingSendStream {
async fn send(
&mut self,
_message: &dyn SendMessage,
_options: SendOptions,
) -> Result<(), ()> {
Err(())
}
}

#[test]
fn grpc_rust_send_encodes_correctly() -> anyhow::Result<()> {
// Arrange
let want = TestMessage {
value: "hello".to_string(),
};
let want = TestMessage::new("hello");

// Act
let mut encoded = GrpcRustSend(want.clone())
Expand All @@ -251,10 +220,8 @@ mod tests {
#[tokio::test]
async fn send_task_join_returns_status_when_send_fails() -> anyhow::Result<()> {
// Arrange
let stream = tokio_stream::iter(vec![TestMessage {
value: "hello".to_string(),
}]);
let mut task = SendTask::start(TestFailingSendStream, stream);
let stream = tokio_stream::iter(vec![TestMessage::new("hello")]);
let mut task = SendTask::start(FailingSendStream, stream);

// Act
let status = task
Expand All @@ -273,22 +240,8 @@ mod tests {
// Arrange
const SIMULATED_PANIC_MSG: &str = "simulated panic in send stream";

struct TestPanicSendStream;

impl SendStream for TestPanicSendStream {
async fn send(
&mut self,
_message: &dyn SendMessage,
_options: SendOptions,
) -> Result<(), ()> {
panic!("{SIMULATED_PANIC_MSG}");
}
}

let stream = tokio_stream::iter(vec![TestMessage {
value: "hello".to_string(),
}]);
let mut task = SendTask::start(TestPanicSendStream, stream);
let stream = tokio_stream::iter(vec![TestMessage::new("hello")]);
let mut task = SendTask::start(PanicSendStream::new(SIMULATED_PANIC_MSG), stream);

// Act
let status = task
Expand All @@ -306,7 +259,7 @@ mod tests {
async fn send_task_join_clears_handle_allowing_safe_drop() -> anyhow::Result<()> {
// Arrange
let stream = tokio_stream::empty::<TestMessage>();
let mut task = SendTask::start(TestPendingSendStream, stream);
let mut task = SendTask::start(PendingSendStream, stream);

// Act
let result = task.join().await;
Expand All @@ -322,7 +275,7 @@ mod tests {
async fn send_task_abort_cancels_task_and_clears_handle() -> anyhow::Result<()> {
// Arrange
let stream = tokio_stream::pending::<TestMessage>();
let mut task = SendTask::start(TestPendingSendStream, stream);
let mut task = SendTask::start(PendingSendStream, stream);
assert!(
task.handle.is_some(),
"task should contain handle after start"
Expand All @@ -343,7 +296,7 @@ mod tests {
async fn send_task_abort_is_idempotent_when_handle_is_none() -> anyhow::Result<()> {
// Arrange
let stream = tokio_stream::pending::<TestMessage>();
let mut task = SendTask::start(TestPendingSendStream, stream);
let mut task = SendTask::start(PendingSendStream, stream);

// Act
task.abort();
Expand All @@ -359,7 +312,7 @@ mod tests {
async fn send_state_join_transitions_to_complete_on_success() -> anyhow::Result<()> {
// Arrange
let stream = tokio_stream::empty::<TestMessage>();
let mut state = SendState::new(SendTask::start(TestPendingSendStream, stream));
let mut state = SendState::new(SendTask::start(PendingSendStream, stream));
assert!(state.is_active(), "state should initially be active");

// Act
Expand All @@ -374,10 +327,8 @@ mod tests {
#[tokio::test]
async fn send_state_join_transitions_to_failed_on_send_error() -> anyhow::Result<()> {
// Arrange
let stream = tokio_stream::iter([TestMessage {
value: "hello".to_string(),
}]);
let mut state = SendState::new(SendTask::start(TestFailingSendStream, stream));
let stream = tokio_stream::iter([TestMessage::new("hello")]);
let mut state = SendState::new(SendTask::start(FailingSendStream, stream));

// Act
state.join().await;
Expand All @@ -393,10 +344,8 @@ mod tests {
#[tokio::test]
async fn send_state_join_if_finished_joins_when_task_is_finished() -> anyhow::Result<()> {
// Arrange
let stream = tokio_stream::iter([TestMessage {
value: "hello".to_string(),
}]);
let mut state = SendState::new(SendTask::start(TestFailingSendStream, stream));
let stream = tokio_stream::iter([TestMessage::new("hello")]);
let mut state = SendState::new(SendTask::start(FailingSendStream, stream));

// Wait briefly for the task to finish executing.
tokio::time::sleep(std::time::Duration::from_millis(50)).await;
Expand All @@ -416,7 +365,7 @@ mod tests {
async fn send_state_join_if_finished_noop_when_task_is_not_finished() -> anyhow::Result<()> {
// Arrange
let stream = tokio_stream::pending::<TestMessage>();
let mut state = SendState::new(SendTask::start(TestPendingSendStream, stream));
let mut state = SendState::new(SendTask::start(PendingSendStream, stream));

// Act
state.join_if_finished().await;
Expand All @@ -431,7 +380,7 @@ mod tests {
async fn send_state_abort_transitions_to_complete_and_aborts_task() -> anyhow::Result<()> {
// Arrange
let stream = tokio_stream::pending::<TestMessage>();
let mut state = SendState::new(SendTask::start(TestPendingSendStream, stream));
let mut state = SendState::new(SendTask::start(PendingSendStream, stream));

// Act
state.abort();
Expand All @@ -451,10 +400,8 @@ mod tests {
#[tokio::test]
async fn send_state_abort_on_failed_state_retains_failure() -> anyhow::Result<()> {
// Arrange
let stream = tokio_stream::iter([TestMessage {
value: "hello".to_string(),
}]);
let mut state = SendState::new(SendTask::start(TestFailingSendStream, stream));
let stream = tokio_stream::iter([TestMessage::new("hello")]);
let mut state = SendState::new(SendTask::start(FailingSendStream, stream));
state.join().await;
assert!(state.failure().is_some(), "should be failed");

Expand Down
Loading
Loading