diff --git a/src/gax-internal/src/grpc/grpc_rust.rs b/src/gax-internal/src/grpc/grpc_rust.rs index 28ecb891af..ca9a831926 100644 --- a/src/gax-internal/src/grpc/grpc_rust.rs +++ b/src/gax-internal/src/grpc/grpc_rust.rs @@ -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; diff --git a/src/gax-internal/src/grpc/grpc_rust/receive.rs b/src/gax-internal/src/grpc/grpc_rust/receive.rs index 5384db858c..cb46528393 100644 --- a/src/gax-internal/src/grpc/grpc_rust/receive.rs +++ b/src/gax-internal/src/grpc/grpc_rust/receive.rs @@ -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; @@ -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::::new(); @@ -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::(TestClosedStream); + let (mut rx, mut task) = ReceiveTask::start::(ClosedRecvStream); // Act let item = rx @@ -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::(TestPendingStream); + let (rx, mut task) = ReceiveTask::start::(PendingRecvStream); drop(rx); // Act @@ -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::(TestPanicStream); + let (_rx, mut task) = + ReceiveTask::start::(PanicRecvStream::new(SIMULATED_PANIC_MSG)); // Act let status = task.join().await; @@ -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::(TestPendingStream); + let (_rx, mut task) = ReceiveTask::start::(PendingRecvStream); let handle = task .handle .as_ref() @@ -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::(TestClosedStream); + let (_rx, mut task) = ReceiveTask::start::(ClosedRecvStream); // Act let status = task.join().await; @@ -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::(TestPendingStream); + let (_rx, mut task) = ReceiveTask::start::(PendingRecvStream); let handle = task .handle .as_ref() @@ -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::(TestHeadersStream { - state: StreamState::default(), - }); + let (mut rx, _task) = ReceiveTask::start::(stream); // Act let first_item = rx @@ -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 diff --git a/src/gax-internal/src/grpc/grpc_rust/send.rs b/src/gax-internal/src/grpc/grpc_rust/send.rs index a0acfd4495..b67cda6264 100644 --- a/src/gax-internal/src/grpc/grpc_rust/send.rs +++ b/src/gax-internal/src/grpc/grpc_rust/send.rs @@ -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()) @@ -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 @@ -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 @@ -306,7 +259,7 @@ mod tests { async fn send_task_join_clears_handle_allowing_safe_drop() -> anyhow::Result<()> { // Arrange let stream = tokio_stream::empty::(); - let mut task = SendTask::start(TestPendingSendStream, stream); + let mut task = SendTask::start(PendingSendStream, stream); // Act let result = task.join().await; @@ -322,7 +275,7 @@ mod tests { async fn send_task_abort_cancels_task_and_clears_handle() -> anyhow::Result<()> { // Arrange let stream = tokio_stream::pending::(); - 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" @@ -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::(); - let mut task = SendTask::start(TestPendingSendStream, stream); + let mut task = SendTask::start(PendingSendStream, stream); // Act task.abort(); @@ -359,7 +312,7 @@ mod tests { async fn send_state_join_transitions_to_complete_on_success() -> anyhow::Result<()> { // Arrange let stream = tokio_stream::empty::(); - 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 @@ -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; @@ -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; @@ -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::(); - let mut state = SendState::new(SendTask::start(TestPendingSendStream, stream)); + let mut state = SendState::new(SendTask::start(PendingSendStream, stream)); // Act state.join_if_finished().await; @@ -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::(); - let mut state = SendState::new(SendTask::start(TestPendingSendStream, stream)); + let mut state = SendState::new(SendTask::start(PendingSendStream, stream)); // Act state.abort(); @@ -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"); diff --git a/src/gax-internal/src/grpc/grpc_rust/testing.rs b/src/gax-internal/src/grpc/grpc_rust/testing.rs new file mode 100644 index 0000000000..3eacff05af --- /dev/null +++ b/src/gax-internal/src/grpc/grpc_rust/testing.rs @@ -0,0 +1,158 @@ +// Copyright 2026 Google LLC +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// https://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +//! Common test helpers, fixtures, and mock streams for `grpc_rust` unit tests. + +use bytes::Bytes; +use grpc::client::{ + RecvStream, ResponseHeaders, ResponseStreamItem, SendOptions, SendStream, Trailers, +}; +use grpc::core::{RecvMessage, SendMessage}; +use prost::Message; +use std::collections::VecDeque; + +#[derive(Clone, PartialEq, Message)] +pub struct TestMessage { + /// The string payload value. + #[prost(string, tag = "1")] + pub value: String, +} + +impl TestMessage { + /// Creates a new [`TestMessage`] with the provided value. + pub fn new(value: impl Into) -> Self { + Self { + value: value.into(), + } + } +} + +/// A [`SendStream`] that immediately fails every send attempt. +pub struct FailingSendStream; + +impl SendStream for FailingSendStream { + async fn send(&mut self, _message: &dyn SendMessage, _options: SendOptions) -> Result<(), ()> { + Err(()) + } +} + +/// A [`SendStream`] that returns pending forever. +pub struct PendingSendStream; + +impl SendStream for PendingSendStream { + async fn send(&mut self, _message: &dyn SendMessage, _options: SendOptions) -> Result<(), ()> { + std::future::pending().await + } +} + +/// A [`SendStream`] that panics on send with a given message. +pub struct PanicSendStream { + /// Panic message to emit. + pub panic_msg: &'static str, +} + +impl PanicSendStream { + /// Creates a new [`PanicSendStream`]. + pub const fn new(panic_msg: &'static str) -> Self { + Self { panic_msg } + } +} + +impl SendStream for PanicSendStream { + async fn send(&mut self, _message: &dyn SendMessage, _options: SendOptions) -> Result<(), ()> { + panic!("{}", self.panic_msg); + } +} + +/// Scripted actions yielded by [`MockRecvStream`]. +#[non_exhaustive] +pub enum MockRecvAction { + /// Yield response headers. + Headers(ResponseHeaders), + /// Yield a response message. + Message(TestMessage), + /// Yield response trailers. + Trailers(Trailers), +} + +/// A mock [`RecvStream`] that executes a queue of [`MockRecvAction`]s. +pub struct MockRecvStream { + actions: VecDeque, +} + +impl MockRecvStream { + /// Creates a new [`MockRecvStream`] from an iterator of actions. + pub fn new(actions: impl IntoIterator) -> Self { + Self { + actions: actions.into_iter().collect(), + } + } +} + +impl RecvStream for MockRecvStream { + async fn recv(&mut self, message: &mut dyn RecvMessage) -> ResponseStreamItem { + if let Some(action) = self.actions.pop_front() { + match action { + MockRecvAction::Headers(headers) => ResponseStreamItem::Headers(headers), + MockRecvAction::Message(response) => { + let mut encoded = Bytes::from(response.encode_to_vec()); + message + .decode(&mut encoded) + .expect("decode response message"); + ResponseStreamItem::Message + } + MockRecvAction::Trailers(trailers) => ResponseStreamItem::Trailers(trailers), + } + } else { + ResponseStreamItem::StreamClosed + } + } +} + +/// A [`RecvStream`] that is immediately closed. +pub struct ClosedRecvStream; + +impl RecvStream for ClosedRecvStream { + async fn recv(&mut self, _message: &mut dyn RecvMessage) -> ResponseStreamItem { + ResponseStreamItem::StreamClosed + } +} + +/// A [`RecvStream`] that stays pending forever. +pub struct PendingRecvStream; + +impl RecvStream for PendingRecvStream { + async fn recv(&mut self, _message: &mut dyn RecvMessage) -> ResponseStreamItem { + std::future::pending().await + } +} + +/// A [`RecvStream`] that panics when polled. +pub struct PanicRecvStream { + /// Panic message to emit. + pub panic_msg: &'static str, +} + +impl PanicRecvStream { + /// Creates a new [`PanicRecvStream`]. + pub const fn new(panic_msg: &'static str) -> Self { + Self { panic_msg } + } +} + +impl RecvStream for PanicRecvStream { + async fn recv(&mut self, _message: &mut dyn RecvMessage) -> ResponseStreamItem { + panic!("{}", self.panic_msg); + } +}