From 05e489901b66a5efdb84db8986cce5b37fe751fe Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Knut=20Olav=20L=C3=B8ite?= Date: Tue, 8 Sep 2026 16:36:39 +0200 Subject: [PATCH] chore(spanner): add stream lifetime guards and error tracking to server streaming --- src/spanner/src/server_streaming/builder.rs | 427 ++++++++++++++++++-- src/spanner/src/server_streaming/stream.rs | 124 ++++-- 2 files changed, 483 insertions(+), 68 deletions(-) diff --git a/src/spanner/src/server_streaming/builder.rs b/src/spanner/src/server_streaming/builder.rs index 154c021469..211de56bf3 100644 --- a/src/spanner/src/server_streaming/builder.rs +++ b/src/spanner/src/server_streaming/builder.rs @@ -23,7 +23,8 @@ use crate::model::ReadRequest; use crate::server_streaming::stream::BatchWriteStream; use crate::server_streaming::stream::CacheUpdateStream; use crate::server_streaming::stream::PartialResultSetStream; -use gaxi::grpc::tonic; +use crate::server_streaming::stream::SpannerServerStream; +use crate::server_streaming::stream::StreamLifetimeGuard; use gaxi::grpc::tonic::Extensions; use gaxi::grpc::tonic::GrpcMethod; use gaxi::prost::ToProto; @@ -36,6 +37,7 @@ pub(crate) struct ExecuteStreamingSql { grpc_client: gaxi::grpc::Client, request: ExecuteSqlRequest, options: RequestOptions, + lifetime_guard: Option, } impl ExecuteStreamingSql { @@ -44,9 +46,17 @@ impl ExecuteStreamingSql { grpc_client, request: ExecuteSqlRequest::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(); @@ -59,23 +69,26 @@ 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( + 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)) + .await } } @@ -91,6 +104,7 @@ pub(crate) struct StreamingRead { grpc_client: gaxi::grpc::Client, request: ReadRequest, options: RequestOptions, + lifetime_guard: Option, } impl StreamingRead { @@ -99,9 +113,17 @@ impl StreamingRead { grpc_client, request: ReadRequest::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(); @@ -114,23 +136,26 @@ 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( + 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)) + .await } } @@ -146,6 +171,7 @@ pub(crate) struct BatchWrite { grpc_client: gaxi::grpc::Client, request: BatchWriteRequest, options: RequestOptions, + lifetime_guard: Option, } impl BatchWrite { @@ -154,9 +180,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(); @@ -171,21 +205,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 } } @@ -201,6 +232,7 @@ pub(crate) struct FetchCacheUpdate { grpc_client: gaxi::grpc::Client, request: FetchCacheUpdateRequest, options: RequestOptions, + lifetime_guard: Option, } impl FetchCacheUpdate { @@ -209,9 +241,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(); @@ -226,21 +266,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 } } @@ -266,20 +303,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, @@ -289,21 +327,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: Clone, Debug, Send, Sync); + static_assertions::assert_impl_all!(StreamingRead: Clone, 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()) @@ -315,18 +394,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 76b0c557da..6e49a139d9 100644 --- a/src/spanner/src/server_streaming/stream.rs +++ b/src/spanner/src/server_streaming/stream.rs @@ -17,11 +17,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; +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)] @@ -32,21 +39,18 @@ 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, } } - /// 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 - } - /// Returns the initial response headers for the stream. pub(crate) fn headers(&self) -> &HeaderMap { &self.headers @@ -60,9 +64,17 @@ impl SpannerServerStream { pub(crate) async fn next_message(&mut self) -> Option> { match self.inner.message().await.map_err(to_gax_error).transpose() { Some(Ok(message)) => Some(Ok(message)), - other => { + Some(Err(err)) => { + if let Some(guard) = self.lifetime_guard.take() + && let Some(status) = err.status() + { + guard.record_error_code(status.code); + } + Some(Err(err)) + } + None => { self.lifetime_guard = None; - other + None } } } @@ -84,8 +96,10 @@ mod tests { 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::atomic::{AtomicBool, Ordering}; + use std::sync::{ + Arc, Mutex, + atomic::{AtomicBool, Ordering}, + }; #[test] fn auto_traits() { @@ -94,8 +108,16 @@ mod tests { static_assertions::assert_impl_all!(CacheUpdateStream: Send, Sync, Debug); } + #[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 { @@ -107,8 +129,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(); @@ -124,9 +148,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), @@ -139,14 +163,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(); @@ -161,9 +192,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); @@ -179,14 +210,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(); @@ -201,9 +239,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"))) @@ -225,6 +263,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(()) } }