diff --git a/src/spanner/src/server_streaming/builder.rs b/src/spanner/src/server_streaming/builder.rs index 3db3791b13..006a2b42b4 100644 --- a/src/spanner/src/server_streaming/builder.rs +++ b/src/spanner/src/server_streaming/builder.rs @@ -23,8 +23,9 @@ use crate::model::ReadRequest; use crate::server_streaming::stream::BatchWriteStream; use crate::server_streaming::stream::CacheUpdateStream; use crate::server_streaming::stream::PartialResultSetStream; +use crate::server_streaming::stream::SpannerServerStream; +use crate::server_streaming::stream::StreamLifetimeGuard; use crate::server_streaming::stream::TransactionIdCallback; -use gaxi::grpc::tonic; use gaxi::grpc::tonic::Extensions; use gaxi::grpc::tonic::GrpcMethod; use gaxi::prost::ToProto; @@ -37,6 +38,7 @@ pub(crate) struct ExecuteStreamingSql { grpc_client: gaxi::grpc::Client, request: ExecuteSqlRequest, options: RequestOptions, + lifetime_guard: Option, on_first_transaction_id: Option, } @@ -46,10 +48,18 @@ impl ExecuteStreamingSql { grpc_client, request: ExecuteSqlRequest::default(), options: RequestOptions::default(), + lifetime_guard: None, on_first_transaction_id: None, } } + /// Attaches an opaque RAII lifetime guard that remains alive for the duration of the stream. + #[allow(dead_code)] + pub(crate) fn with_lifetime_guard(mut self, guard: StreamLifetimeGuard) -> Self { + self.lifetime_guard = Some(guard); + self + } + /// Sets the full request, replacing any prior values. pub(crate) fn with_request>(mut self, v: V) -> Self { self.request = v.into(); @@ -71,24 +81,27 @@ impl ExecuteStreamingSql { self } + /// Returns a reference to the request options. + #[allow(dead_code)] + pub(crate) fn options(&self) -> &RequestOptions { + &self.options + } + /// Start the server streaming request and receive the stream. pub(crate) async fn send(self) -> Result { - let session = self.request.session.clone(); + let request_params = format!("session={}", self.request.session); let request = self.request.to_proto().map_err(Error::deser)?; - let request_params = format!("session={session}"); - let response = make_server_streaming_request( + let stream = make_server_streaming_request( &self.grpc_client, request, self.options, "ExecuteStreamingSql", "/google.spanner.v1.Spanner/ExecuteStreamingSql", &request_params, + self.lifetime_guard, ) .await?; - let (metadata, stream, _) = response.into_parts(); - let headers = metadata.into_headers(); - Ok(PartialResultSetStream::new(stream, headers) - .with_transaction_id_callback(self.on_first_transaction_id)) + Ok(stream.with_transaction_id_callback(self.on_first_transaction_id)) } } @@ -104,6 +117,7 @@ pub(crate) struct StreamingRead { grpc_client: gaxi::grpc::Client, request: ReadRequest, options: RequestOptions, + lifetime_guard: Option, on_first_transaction_id: Option, } @@ -113,10 +127,18 @@ impl StreamingRead { grpc_client, request: ReadRequest::default(), options: RequestOptions::default(), + lifetime_guard: None, on_first_transaction_id: None, } } + /// Attaches an opaque RAII lifetime guard that remains alive for the duration of the stream. + #[allow(dead_code)] + pub(crate) fn with_lifetime_guard(mut self, guard: StreamLifetimeGuard) -> Self { + self.lifetime_guard = Some(guard); + self + } + /// Sets the full request, replacing any prior values. pub(crate) fn with_request>(mut self, v: V) -> Self { self.request = v.into(); @@ -138,24 +160,27 @@ impl StreamingRead { self } + /// Returns a reference to the request options. + #[allow(dead_code)] + pub(crate) fn options(&self) -> &RequestOptions { + &self.options + } + /// Start the server streaming request and receive the stream. pub(crate) async fn send(self) -> Result { - let session = self.request.session.clone(); + let request_params = format!("session={}", self.request.session); let request = self.request.to_proto().map_err(Error::deser)?; - let request_params = format!("session={session}"); - let response = make_server_streaming_request( + let stream = make_server_streaming_request( &self.grpc_client, request, self.options, "StreamingRead", "/google.spanner.v1.Spanner/StreamingRead", &request_params, + self.lifetime_guard, ) .await?; - let (metadata, stream, _) = response.into_parts(); - let headers = metadata.into_headers(); - Ok(PartialResultSetStream::new(stream, headers) - .with_transaction_id_callback(self.on_first_transaction_id)) + Ok(stream.with_transaction_id_callback(self.on_first_transaction_id)) } } @@ -171,6 +196,7 @@ pub(crate) struct BatchWrite { grpc_client: gaxi::grpc::Client, request: BatchWriteRequest, options: RequestOptions, + lifetime_guard: Option, } impl BatchWrite { @@ -179,9 +205,17 @@ impl BatchWrite { grpc_client, request: BatchWriteRequest::default(), options: RequestOptions::default(), + lifetime_guard: None, } } + /// Attaches an opaque RAII lifetime guard that remains alive for the duration of the stream. + #[allow(dead_code)] + pub(crate) fn with_lifetime_guard(mut self, guard: StreamLifetimeGuard) -> Self { + self.lifetime_guard = Some(guard); + self + } + /// Sets the full request, replacing any prior values. pub(crate) fn with_request>(mut self, v: V) -> Self { self.request = v.into(); @@ -196,21 +230,18 @@ impl BatchWrite { /// Start the server streaming request and receive the stream. pub(crate) async fn send(self) -> Result { - let session = self.request.session.clone(); + let request_params = format!("session={}", self.request.session); let request = self.request.to_proto().map_err(Error::deser)?; - let request_params = format!("session={session}"); - let response = make_server_streaming_request( + make_server_streaming_request( &self.grpc_client, request, self.options, "BatchWrite", "/google.spanner.v1.Spanner/BatchWrite", &request_params, + self.lifetime_guard, ) - .await?; - let (metadata, stream, _) = response.into_parts(); - let headers = metadata.into_headers(); - Ok(BatchWriteStream::new(stream, headers)) + .await } } @@ -226,6 +257,7 @@ pub(crate) struct FetchCacheUpdate { grpc_client: gaxi::grpc::Client, request: FetchCacheUpdateRequest, options: RequestOptions, + lifetime_guard: Option, } impl FetchCacheUpdate { @@ -234,9 +266,17 @@ impl FetchCacheUpdate { grpc_client, request: FetchCacheUpdateRequest::default(), options: RequestOptions::default(), + lifetime_guard: None, } } + /// Attaches an opaque RAII lifetime guard that remains alive for the duration of the stream. + #[allow(dead_code)] + pub(crate) fn with_lifetime_guard(mut self, guard: StreamLifetimeGuard) -> Self { + self.lifetime_guard = Some(guard); + self + } + /// Sets the full request, replacing any prior values. pub(crate) fn with_request>(mut self, v: V) -> Self { self.request = v.into(); @@ -251,21 +291,18 @@ impl FetchCacheUpdate { /// Start the server streaming request and receive the stream. pub(crate) async fn send(self) -> Result { - let database = self.request.database.clone(); + let request_params = format!("database={}", self.request.database); let request = self.request.to_proto().map_err(Error::deser)?; - let request_params = format!("database={database}"); - let response = make_server_streaming_request( + make_server_streaming_request( &self.grpc_client, request, self.options, "FetchCacheUpdate", "/google.spanner.v1.Spanner/FetchCacheUpdate", &request_params, + self.lifetime_guard, ) - .await?; - let (metadata, stream, _) = response.into_parts(); - let headers = metadata.into_headers(); - Ok(CacheUpdateStream::new(stream, headers)) + .await } } @@ -291,20 +328,21 @@ async fn make_server_streaming_request( method_name: &'static str, path_str: &'static str, x_goog_request_params: &str, -) -> Result>> + lifetime_guard: Option, +) -> Result> where Req: Message + Default + Clone + 'static, Res: Message + Default + 'static, { let options = google_cloud_gax::options::internal::set_default_idempotency(options, false); let extensions = { - let mut e = Extensions::new(); - e.insert(GrpcMethod::new("google.spanner.v1.Spanner", method_name)); - e + let mut extensions = Extensions::new(); + extensions.insert(GrpcMethod::new("google.spanner.v1.Spanner", method_name)); + extensions }; let path = http::uri::PathAndQuery::from_static(path_str); - grpc_client + let response = match grpc_client .server_streaming( extensions, path, @@ -314,21 +352,62 @@ where x_goog_request_params, ) .await + { + Ok(response) => response, + Err(err) => { + if let Some(guard) = lifetime_guard + && let Some(status) = err.status() + { + guard.record_error_code(status.code); + } + return Err(err); + } + }; + let (metadata, stream, _) = response.into_parts(); + let headers = metadata.into_headers(); + Ok(SpannerServerStream::new(stream, headers, lifetime_guard)) } #[cfg(test)] mod tests { use super::*; use crate::client::Spanner; + use crate::server_streaming::stream::StreamGuard; + use gaxi::grpc::tonic::Status; use google_cloud_auth::credentials::anonymous::Builder as Anonymous; + use google_cloud_gax::error::rpc::Code; + use google_cloud_gax::options::RequestOptions; use google_cloud_test_macros::tokio_test_no_panics; + use std::fmt::Debug; + use std::sync::Arc; + use std::sync::Mutex; + use tokio::task::JoinHandle; + + #[test] + fn traits() { + static_assertions::assert_impl_all!(ExecuteStreamingSql: Debug, Send, Sync); + static_assertions::assert_impl_all!(StreamingRead: Debug, Send, Sync); + static_assertions::assert_impl_all!(BatchWrite: Clone, Debug, Send, Sync); + static_assertions::assert_impl_all!(FetchCacheUpdate: Clone, Debug, Send, Sync); + } - #[tokio_test_no_panics] - async fn fetch_cache_update_builder_configuration() { - let (address, _server) = - spanner_grpc_mock::start("0.0.0.0:0", spanner_grpc_mock::MockSpanner::new()) - .await - .expect("mock server should start"); + #[derive(Debug)] + struct TestLifetimeGuard { + recorded_code: Arc>>, + } + + impl StreamGuard for TestLifetimeGuard { + fn record_error_code(&self, code: Code) { + *self.recorded_code.lock().expect("lock poisoned") = Some(code); + } + } + + async fn create_test_grpc_client_with_mock( + mock: spanner_grpc_mock::MockSpanner, + ) -> (gaxi::grpc::Client, JoinHandle<()>) { + let (address, server) = spanner_grpc_mock::start("0.0.0.0:0", mock) + .await + .expect("mock server should start"); let spanner = Spanner::builder() .with_endpoint(address) .with_credentials(Anonymous::new().build()) @@ -340,18 +419,280 @@ mod tests { .grpc_client .clone() .expect("grpc client should exist"); + (grpc_client, server) + } + + async fn create_test_grpc_client() -> (gaxi::grpc::Client, JoinHandle<()>) { + create_test_grpc_client_with_mock(spanner_grpc_mock::MockSpanner::new()).await + } + + #[tokio_test_no_panics] + async fn execute_streaming_sql_builder_configuration() { + let (grpc_client, _server) = create_test_grpc_client().await; + let guard = Arc::new(TestLifetimeGuard { + recorded_code: Arc::new(Mutex::new(None)), + }); + + let mut options = RequestOptions::default(); + options.set_idempotency(true); + + let mut builder = ExecuteStreamingSql::new(grpc_client) + .with_request( + ExecuteSqlRequest::default() + .set_session("projects/p/instances/i/databases/d/sessions/s") + .set_sql("SELECT 1"), + ) + .with_options(options) + .with_lifetime_guard(guard); + + assert_eq!( + builder.options().idempotent(), + Some(true), + "builder options should reflect configured options" + ); + let request_options = builder.request_options(); + request_options.set_idempotency(false); + assert_eq!( + builder.options().idempotent(), + Some(false), + "builder options should reflect mutated options" + ); + assert_eq!( + builder.request.session, "projects/p/instances/i/databases/d/sessions/s", + "session should match configured value" + ); + assert_eq!( + builder.request.sql, "SELECT 1", + "sql should match configured value" + ); + } + + #[tokio_test_no_panics] + async fn streaming_read_builder_configuration() { + let (grpc_client, _server) = create_test_grpc_client().await; + let guard = Arc::new(TestLifetimeGuard { + recorded_code: Arc::new(Mutex::new(None)), + }); + + let mut options = RequestOptions::default(); + options.set_idempotency(true); + + let mut builder = StreamingRead::new(grpc_client) + .with_request( + ReadRequest::default() + .set_session("projects/p/instances/i/databases/d/sessions/s") + .set_table("MyTable"), + ) + .with_options(options) + .with_lifetime_guard(guard); + + assert_eq!( + builder.options().idempotent(), + Some(true), + "builder options should reflect configured options" + ); + let request_options = builder.request_options(); + request_options.set_idempotency(false); + assert_eq!( + builder.options().idempotent(), + Some(false), + "builder options should reflect mutated options" + ); + assert_eq!( + builder.request.session, "projects/p/instances/i/databases/d/sessions/s", + "session should match configured value" + ); + assert_eq!( + builder.request.table, "MyTable", + "table should match configured value" + ); + } + + #[tokio_test_no_panics] + async fn batch_write_builder_configuration() { + let (grpc_client, _server) = create_test_grpc_client().await; + let guard = Arc::new(TestLifetimeGuard { + recorded_code: Arc::new(Mutex::new(None)), + }); + + let mut builder = BatchWrite::new(grpc_client) + .with_request( + BatchWriteRequest::default() + .set_session("projects/p/instances/i/databases/d/sessions/s"), + ) + .with_options(RequestOptions::default()) + .with_lifetime_guard(guard); + + let _ = builder.request_options(); + assert_eq!( + builder.request.session, "projects/p/instances/i/databases/d/sessions/s", + "session should match configured value" + ); + } + + #[tokio_test_no_panics] + async fn fetch_cache_update_builder_configuration() { + let (grpc_client, _server) = create_test_grpc_client().await; + let guard = Arc::new(TestLifetimeGuard { + recorded_code: Arc::new(Mutex::new(None)), + }); let mut builder = FetchCacheUpdate::new(grpc_client) .with_request( FetchCacheUpdateRequest::default() .set_database("projects/p/instances/i/databases/d"), ) - .with_options(RequestOptions::default()); + .with_options(RequestOptions::default()) + .with_lifetime_guard(guard); let _ = builder.request_options(); assert_eq!( - builder.request.database, - "projects/p/instances/i/databases/d" + builder.request.database, "projects/p/instances/i/databases/d", + "database should match configured value" + ); + } + + #[tokio_test_no_panics] + async fn make_server_streaming_request_records_error_code_on_initial_failure() { + let mut mock = spanner_grpc_mock::MockSpanner::new(); + mock.expect_execute_streaming_sql() + .once() + .returning(|_| Err(Status::unavailable("backend unavailable"))); + + let (grpc_client, _server) = create_test_grpc_client_with_mock(mock).await; + + let recorded_code = Arc::new(Mutex::new(None)); + let guard = Arc::new(TestLifetimeGuard { + recorded_code: Arc::clone(&recorded_code), + }); + + let builder = ExecuteStreamingSql::new(grpc_client) + .with_lifetime_guard(guard) + .with_request( + ExecuteSqlRequest::default() + .set_session("projects/p/instances/i/databases/d/sessions/s") + .set_sql("SELECT 1"), + ); + + let result = builder.send().await; + assert!(result.is_err(), "Initial handshake failure must return Err"); + assert_eq!( + *recorded_code.lock().expect("lock poisoned"), + Some(Code::Unavailable), + "Initial handshake failure must record Code::Unavailable on lifetime guard" + ); + } + + #[tokio_test_no_panics] + async fn streaming_read_records_error_code_on_initial_failure() { + let mut mock = spanner_grpc_mock::MockSpanner::new(); + mock.expect_streaming_read() + .once() + .returning(|_| Err(Status::unavailable("backend unavailable"))); + + let (grpc_client, _server) = create_test_grpc_client_with_mock(mock).await; + + let recorded_code = Arc::new(Mutex::new(None)); + let guard = Arc::new(TestLifetimeGuard { + recorded_code: Arc::clone(&recorded_code), + }); + + let builder = StreamingRead::new(grpc_client) + .with_lifetime_guard(guard) + .with_request( + ReadRequest::default() + .set_session("projects/p/instances/i/databases/d/sessions/s") + .set_table("MyTable"), + ); + + let result = builder.send().await; + assert!(result.is_err(), "Initial handshake failure must return Err"); + assert_eq!( + *recorded_code.lock().expect("lock poisoned"), + Some(Code::Unavailable), + "Initial handshake failure must record Code::Unavailable on lifetime guard" + ); + } + + #[tokio_test_no_panics] + async fn batch_write_records_error_code_on_initial_failure() { + let mut mock = spanner_grpc_mock::MockSpanner::new(); + mock.expect_batch_write() + .once() + .returning(|_| Err(Status::unavailable("backend unavailable"))); + + let (grpc_client, _server) = create_test_grpc_client_with_mock(mock).await; + + let recorded_code = Arc::new(Mutex::new(None)); + let guard = Arc::new(TestLifetimeGuard { + recorded_code: Arc::clone(&recorded_code), + }); + + let builder = BatchWrite::new(grpc_client) + .with_lifetime_guard(guard) + .with_request( + BatchWriteRequest::default() + .set_session("projects/p/instances/i/databases/d/sessions/s"), + ); + + let result = builder.send().await; + assert!(result.is_err(), "Initial handshake failure must return Err"); + assert_eq!( + *recorded_code.lock().expect("lock poisoned"), + Some(Code::Unavailable), + "Initial handshake failure must record Code::Unavailable on lifetime guard" + ); + } + + #[tokio_test_no_panics] + async fn fetch_cache_update_records_error_code_on_initial_failure() { + let mut mock = spanner_grpc_mock::MockSpanner::new(); + mock.expect_fetch_cache_update() + .once() + .returning(|_| Err(Status::unavailable("backend unavailable"))); + + let (grpc_client, _server) = create_test_grpc_client_with_mock(mock).await; + + let recorded_code = Arc::new(Mutex::new(None)); + let guard = Arc::new(TestLifetimeGuard { + recorded_code: Arc::clone(&recorded_code), + }); + + let builder = FetchCacheUpdate::new(grpc_client) + .with_lifetime_guard(guard) + .with_request( + FetchCacheUpdateRequest::default() + .set_database("projects/p/instances/i/databases/d"), + ); + + let result = builder.send().await; + assert!(result.is_err(), "Initial handshake failure must return Err"); + assert_eq!( + *recorded_code.lock().expect("lock poisoned"), + Some(Code::Unavailable), + "Initial handshake failure must record Code::Unavailable on lifetime guard" + ); + } + + #[tokio_test_no_panics] + async fn make_server_streaming_request_without_guard_returns_error() { + let mut mock = spanner_grpc_mock::MockSpanner::new(); + mock.expect_execute_streaming_sql() + .once() + .returning(|_| Err(Status::unavailable("backend unavailable"))); + + let (grpc_client, _server) = create_test_grpc_client_with_mock(mock).await; + + let builder = ExecuteStreamingSql::new(grpc_client).with_request( + ExecuteSqlRequest::default() + .set_session("projects/p/instances/i/databases/d/sessions/s") + .set_sql("SELECT 1"), + ); + + let result = builder.send().await; + assert!( + result.is_err(), + "Initial handshake failure without lifetime guard must return Err" ); } } diff --git a/src/spanner/src/server_streaming/stream.rs b/src/spanner/src/server_streaming/stream.rs index 948e1e29f8..08e8d3448c 100644 --- a/src/spanner/src/server_streaming/stream.rs +++ b/src/spanner/src/server_streaming/stream.rs @@ -18,12 +18,18 @@ use crate::google::spanner::v1::CacheUpdate as ProtoCacheUpdate; use crate::google::spanner::v1::PartialResultSet; use gaxi::grpc::from_status::to_gax_error; use gaxi::grpc::tonic::Streaming; +use google_cloud_gax::error::rpc::Code; use http::HeaderMap; -use std::any::Any; use std::fmt::{Debug, Formatter, Result as FmtResult}; +use std::sync::Arc; -/// Type alias for opaque stream lifetime drop guards. -pub(crate) type StreamLifetimeGuard = Box; +/// Trait for stream lifetime drop guards capable of recording RPC error codes for dynamic channel pooling. +pub(crate) trait StreamGuard: Debug + Send + Sync + 'static { + fn record_error_code(&self, code: Code); +} + +/// Type alias for stream lifetime drop guards. +pub(crate) type StreamLifetimeGuard = Arc; /// Generic wrapper around gRPC server-streaming responses with lifetime management. #[derive(Debug)] @@ -35,11 +41,15 @@ pub(crate) struct SpannerServerStream { } impl SpannerServerStream { - pub(crate) fn new(inner: Streaming, headers: HeaderMap) -> Self { + pub(crate) fn new( + inner: Streaming, + headers: HeaderMap, + lifetime_guard: Option, + ) -> Self { Self { inner, headers, - lifetime_guard: None, + lifetime_guard, on_first_transaction_id: None, } } @@ -75,10 +85,19 @@ impl SpannerServerStream { pub(crate) async fn next_message(&mut self) -> Option> { let message = match self.inner.message().await.map_err(to_gax_error).transpose() { Some(Ok(message)) => message, - other => { + Some(Err(err)) => { + if let Some(guard) = self.lifetime_guard.take() + && let Some(status) = err.status() + { + guard.record_error_code(status.code); + } + self.on_first_transaction_id = None; + return Some(Err(err)); + } + None => { self.lifetime_guard = None; self.on_first_transaction_id = None; - return other; + return None; } }; @@ -196,8 +215,16 @@ mod tests { ); } + #[derive(Debug)] struct TestDropGuard { dropped: Arc, + recorded_code: Arc>>, + } + + impl StreamGuard for TestDropGuard { + fn record_error_code(&self, code: Code) { + *self.recorded_code.lock().expect("lock poisoned") = Some(code); + } } impl Drop for TestDropGuard { @@ -209,8 +236,10 @@ mod tests { #[tokio_test_no_panics] async fn stream_drop_releases_lifetime_guard() -> anyhow::Result<()> { let dropped = Arc::new(AtomicBool::new(false)); - let guard = Box::new(TestDropGuard { + let recorded_code = Arc::new(Mutex::new(None)); + let guard = Arc::new(TestDropGuard { dropped: Arc::clone(&dropped), + recorded_code: Arc::clone(&recorded_code), }); let mut mock = create_session_mock(); @@ -226,9 +255,9 @@ mod tests { .set_sql("SELECT 1"); let stream = db_client .execute_streaming_sql(request, RequestOptions::default(), 0) + .with_lifetime_guard(guard) .send() - .await? - .with_lifetime_guard(guard); + .await?; assert!( !dropped.load(Ordering::Relaxed), @@ -241,14 +270,21 @@ mod tests { dropped.load(Ordering::Relaxed), "Guard must be dropped when stream is dropped" ); + assert_eq!( + *recorded_code.lock().expect("lock poisoned"), + None, + "No error code should be recorded on normal stream drop" + ); Ok(()) } #[tokio_test_no_panics] async fn stream_eof_releases_lifetime_guard() -> anyhow::Result<()> { let dropped = Arc::new(AtomicBool::new(false)); - let guard = Box::new(TestDropGuard { + let recorded_code = Arc::new(Mutex::new(None)); + let guard = Arc::new(TestDropGuard { dropped: Arc::clone(&dropped), + recorded_code: Arc::clone(&recorded_code), }); let mut mock = create_session_mock(); @@ -263,9 +299,9 @@ mod tests { .set_sql("SELECT 1"); let mut stream = db_client .execute_streaming_sql(request, RequestOptions::default(), 0) + .with_lifetime_guard(guard) .send() - .await? - .with_lifetime_guard(guard); + .await?; // Close channel to simulate EOF drop(sender); @@ -281,14 +317,21 @@ mod tests { dropped.load(Ordering::Relaxed), "Guard must be dropped immediately on EOF" ); + assert_eq!( + *recorded_code.lock().expect("lock poisoned"), + None, + "No error code should be recorded on normal EOF" + ); Ok(()) } #[tokio_test_no_panics] async fn stream_error_releases_lifetime_guard() -> anyhow::Result<()> { let dropped = Arc::new(AtomicBool::new(false)); - let guard = Box::new(TestDropGuard { + let recorded_code = Arc::new(Mutex::new(None)); + let guard = Arc::new(TestDropGuard { dropped: Arc::clone(&dropped), + recorded_code: Arc::clone(&recorded_code), }); let mut mock = create_session_mock(); @@ -303,9 +346,9 @@ mod tests { .set_sql("SELECT 1"); let mut stream = db_client .execute_streaming_sql(request, RequestOptions::default(), 0) + .with_lifetime_guard(guard) .send() - .await? - .with_lifetime_guard(guard); + .await?; sender .send(Err(Status::unavailable("server unavailable"))) @@ -327,6 +370,42 @@ mod tests { dropped.load(Ordering::Relaxed), "Guard must be dropped immediately on stream error" ); + assert_eq!( + *recorded_code.lock().expect("lock poisoned"), + Some(Code::Unavailable), + "Stream error must record Code::Unavailable on guard" + ); + Ok(()) + } + + #[tokio_test_no_panics] + async fn stream_error_without_guard() -> anyhow::Result<()> { + let mut mock = create_session_mock(); + let (sender, receiver) = tokio::sync::mpsc::channel(1); + mock.expect_execute_streaming_sql() + .return_once(move |_| Ok(Response::from(receiver))); + + let (db_client, _server) = setup_db_client(mock).await; + + let request = crate::model::ExecuteSqlRequest::default() + .set_session(db_client.session_name()) + .set_sql("SELECT 1"); + let mut stream = db_client + .execute_streaming_sql(request, RequestOptions::default(), 0) + .send() + .await?; + + sender + .send(Err(Status::unavailable("server unavailable"))) + .await + .expect("send error"); + + let next = stream.next_message().await; + assert!(next.is_some(), "Stream should yield Some on error"); + assert!( + next.expect("error message").is_err(), + "Stream message should be an error" + ); Ok(()) }