From 39f89f462475d4a2b9070293f9dcd6ab9d014638 Mon Sep 17 00:00:00 2001 From: Darren Bolduc Date: Tue, 8 Sep 2026 12:34:45 -0400 Subject: [PATCH 1/3] impl(bigquery): introduce stream dispatcher --- Cargo.lock | 10 + Cargo.toml | 1 + src/bigquery/Cargo.toml | 1 + src/bigquery/src/write.rs | 2 + src/bigquery/src/write/dispatcher.rs | 262 +++++++++++++++++++++++++++ src/bigquery/src/write/pool.rs | 8 +- 6 files changed, 283 insertions(+), 1 deletion(-) create mode 100644 src/bigquery/src/write/dispatcher.rs diff --git a/Cargo.lock b/Cargo.lock index c5709c1581..9edd561c37 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -57,6 +57,15 @@ version = "1.0.104" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "330a5ed07fa54e4702c9d6c4174f74427fc0ef6e214bbd677ae50a5099946470" +[[package]] +name = "arc-swap" +version = "1.9.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c049c0be4daef0b145cb3555416b3b8ef5b7888a38aea1a3a155801fe7b0810b" +dependencies = [ + "rustversion", +] + [[package]] name = "arrayvec" version = "0.7.8" @@ -2374,6 +2383,7 @@ name = "google-cloud-bigquery" version = "0.16.1-preview" dependencies = [ "anyhow", + "arc-swap", "async-trait", "base64 0.23.1", "bigquery-grpc-mock", diff --git a/Cargo.toml b/Cargo.toml index bd3b6f97e8..9e01527aad 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -424,6 +424,7 @@ inherits = "dev" incremental = false [workspace.dependencies] +arc-swap = { default-features = false, version = "1.9" } arrow = { default-features = false, version = "59.2" } arrow-json = { default-features = false, version = "59.2" } async-trait = { default-features = false, version = "0.1.85" } diff --git a/src/bigquery/Cargo.toml b/src/bigquery/Cargo.toml index a508bdeb62..6ef261f285 100644 --- a/src/bigquery/Cargo.toml +++ b/src/bigquery/Cargo.toml @@ -26,6 +26,7 @@ categories.workspace = true rust-version.workspace = true [dependencies] +arc-swap.workspace = true async-trait.workspace = true base64.workspace = true bytes.workspace = true diff --git a/src/bigquery/src/write.rs b/src/bigquery/src/write.rs index 997039449e..31565e7a29 100644 --- a/src/bigquery/src/write.rs +++ b/src/bigquery/src/write.rs @@ -28,6 +28,8 @@ pub(super) mod client; pub(super) mod client_builder; pub(super) mod error; +#[cfg_attr(not(test), expect(dead_code))] +mod dispatcher; mod entry; #[cfg_attr(not(test), expect(dead_code))] mod pool; diff --git a/src/bigquery/src/write/dispatcher.rs b/src/bigquery/src/write/dispatcher.rs new file mode 100644 index 0000000000..9da9229d44 --- /dev/null +++ b/src/bigquery/src/write/dispatcher.rs @@ -0,0 +1,262 @@ +// 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. + +use super::append_response::{AppendResponse, to_result}; +use super::entry::StreamEntry; +use super::error::{AppendError, AppendResult}; +use super::pool::StreamPool; +use crate::Error; +use crate::model::AppendRowsRequest; +use arc_swap::ArcSwap; +use gaxi::prost::{FromProto, ToProto}; +use std::sync::Arc; + +/// Efficiently dispatches writes to a stream in a stream pool. +/// +/// This struct caches a `StreamEntry` and loads it atomically on each write. +/// +/// On a transient error, the `Dispatcher` notifies the `StreamPool` of the +/// failed `StreamEntry` and receives a new `StreamEntry` to use for future +/// writes. +/// +/// This struct is also responsible for retrying individual writes. +#[derive(Debug)] +pub(crate) struct Dispatcher { + pub(crate) pool: Arc, + pub(crate) entry: ArcSwap, +} + +impl Dispatcher { + /// Creates a new `Dispatcher` for a given `StreamPool`. + pub(crate) fn new(pool: Arc) -> Self { + let stream = pool.get(); + Self { + pool, + entry: ArcSwap::from_pointee(stream), + } + } + + /// Send the write and process the response. + /// + /// Evicts and updates its cached stream on transient errors. + pub(crate) async fn send(&self, req: AppendRowsRequest) -> AppendResult { + let req = req.to_proto().map_err(Error::deser)?; + + let stream = self.entry.load_full(); + let stream_id = stream.id; + + let resp = match stream.send(req).await { + Ok(resp) => Ok(resp), + Err(err) => { + if is_transient_error(&err) { + // Atomically evicts failed_id and returns a new stream for use. + let new_stream = self.pool.evict_and_replace(stream_id); + + // The application can `send()` multiple writes + // concurrently. Only one `send()` should update the cached + // stream on a transient error. + let _ = self.entry.compare_and_swap(&stream, Arc::new(new_stream)); + + // TODO(#6355): implement retries + } + Err(err) + } + }?; + + let resp = resp.cnv().map_err(Error::ser)?; + to_result(resp) + } +} + +pub(crate) fn is_transient_error(err: &AppendError) -> bool { + match err { + AppendError::UnexpectedEndOfStream => true, + // TODO(#6355): classify transient RPC errors + _ => false, + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::write::test::*; + use bigquery_grpc_mock::{MockBigQueryWrite, start}; + use gaxi::grpc::tonic::{Response as TonicResponse, Status as TonicStatus}; + use tokio::sync::mpsc; + use tokio::task::JoinSet; + + fn test_req() -> AppendRowsRequest { + AppendRowsRequest::new() + } + + #[tokio::test] + async fn success() -> anyhow::Result<()> { + let (response_tx, response_rx) = mpsc::channel(10); + let mut mock = MockBigQueryWrite::new(); + mock.expect_append_rows() + .return_once(move |_| Ok(TonicResponse::from(response_rx))); + + let (endpoint, _server) = start("0.0.0.0:0", mock).await?; + let transport = Arc::new(test_transport(endpoint).await?); + let pool = Arc::new(StreamPool::new(transport, 10)); + let dispatcher = Arc::new(Dispatcher::new(pool)); + assert_eq!(dispatcher.entry.load().id, 1); + + let write1 = { + let d = dispatcher.clone(); + tokio::spawn(async move { d.send(test_req()).await }) + }; + let write2 = { + let d = dispatcher.clone(); + tokio::spawn(async move { d.send(test_req()).await }) + }; + + // Respond to the writes + response_tx.send(Ok(convert(&test_response(1)))).await?; + assert_eq!(write1.await??.offset, Some(1)); + + response_tx.send(Ok(convert(&test_response(2)))).await?; + assert_eq!(write2.await??.offset, Some(2)); + + // Verify we are still on the same stream. + assert_eq!(dispatcher.entry.load().id, 1); + + Ok(()) + } + + #[tokio::test] + async fn stream_closed() -> anyhow::Result<()> { + let (response_tx, response_rx) = mpsc::channel(10); + let mut mock = MockBigQueryWrite::new(); + mock.expect_append_rows() + .return_once(move |_| Ok(TonicResponse::from(response_rx))); + + let (endpoint, _server) = start("0.0.0.0:0", mock).await?; + let transport = Arc::new(test_transport(endpoint).await?); + let pool = Arc::new(StreamPool::new(transport, 10)); + let dispatcher = Arc::new(Dispatcher::new(pool)); + assert_eq!(dispatcher.entry.load().id, 1); + + let write = { + let d = dispatcher.clone(); + tokio::spawn(async move { d.send(test_req()).await }) + }; + + // Simulate the stream closing before responding to the request. + drop(response_tx); + + // TODO(#6355) - expect retries. + let err = write.await?.expect_err("should return an error"); + assert!(matches!(err, AppendError::UnexpectedEndOfStream)); + + // We ran into a transient error. We should now have a new stream. + assert_eq!(dispatcher.entry.load().id, 2); + + Ok(()) + } + + #[tokio::test] + async fn permanent_error() -> anyhow::Result<()> { + let (response_tx, response_rx) = mpsc::channel(10); + let mut mock = MockBigQueryWrite::new(); + mock.expect_append_rows() + .return_once(move |_| Ok(TonicResponse::from(response_rx))); + + let (endpoint, _server) = start("0.0.0.0:0", mock).await?; + let transport = Arc::new(test_transport(endpoint).await?); + let pool = Arc::new(StreamPool::new(transport, 10)); + let dispatcher = Arc::new(Dispatcher::new(pool)); + assert_eq!(dispatcher.entry.load().id, 1); + + let write = { + let d = dispatcher.clone(); + tokio::spawn(async move { d.send(test_req()).await }) + }; + + // Simulate a permanent stream error + response_tx + .send(Err(TonicStatus::failed_precondition("fail"))) + .await?; + + let err = write.await?.expect_err("should return an error"); + assert!(matches!(err, AppendError::Rpc { source: _ })); + + Ok(()) + } + + #[tokio::test] + async fn transient_error_evict_contention() -> anyhow::Result<()> { + let (response_tx, response_rx) = mpsc::channel(10); + let mut mock = MockBigQueryWrite::new(); + mock.expect_append_rows() + .return_once(move |_| Ok(TonicResponse::from(response_rx))); + + let (endpoint, _server) = start("0.0.0.0:0", mock).await?; + let transport = Arc::new(test_transport(endpoint).await?); + let pool = Arc::new(StreamPool::new(transport, 10)); + let dispatcher = Arc::new(Dispatcher::new(pool.clone())); + assert_eq!(dispatcher.entry.load().id, 1); + + let mut writes = JoinSet::new(); + for _ in 0..1000 { + let d = dispatcher.clone(); + writes.spawn(async move { d.send(test_req()).await }); + } + + // Simulate the stream closing before responding to the requests. + drop(response_tx); + + while let Some(write) = writes.join_next().await { + let err = write?.expect_err("should return an error"); + assert!(matches!(err, AppendError::UnexpectedEndOfStream)); + } + + // We ran into a transient error. We should now have a new stream. Only + // one of the dispatchers should have evicted the failed stream. + assert_eq!(dispatcher.entry.load().id, 2); + assert_eq!(pool.stream_ids(), [2]); + + Ok(()) + } + + #[tokio::test] + async fn writes_bypass_pool_lock() -> anyhow::Result<()> { + let (response_tx, response_rx) = mpsc::channel(10); + let mut mock = MockBigQueryWrite::new(); + mock.expect_append_rows() + .return_once(move |_| Ok(TonicResponse::from(response_rx))); + + let (endpoint, _server) = start("0.0.0.0:0", mock).await?; + let transport = Arc::new(test_transport(endpoint).await?); + let pool = Arc::new(StreamPool::new(transport, 10)); + let dispatcher = Arc::new(Dispatcher::new(pool.clone())); + assert_eq!(dispatcher.entry.load().id, 1); + + // Acquire the stream pool's lock to simulate a pool scaling event. + let _guard = pool.lock(); + + let write = { + let d = dispatcher.clone(); + tokio::spawn(async move { d.send(test_req()).await }) + }; + + response_tx.send(Ok(convert(&test_response(1)))).await?; + + // Verify the write goes through, even with the pool's lock held. + assert_eq!(write.await??.offset, Some(1)); + assert_eq!(dispatcher.entry.load().id, 1); + + Ok(()) + } +} diff --git a/src/bigquery/src/write/pool.rs b/src/bigquery/src/write/pool.rs index ec3878e4b9..666bd52ace 100644 --- a/src/bigquery/src/write/pool.rs +++ b/src/bigquery/src/write/pool.rs @@ -137,6 +137,7 @@ mod tests { use crate::write::test::*; use bigquery_grpc_mock::{MockBigQueryWrite, start}; use gaxi::grpc::tonic::Response as TonicResponse; + use std::sync::MutexGuard; use test_case::test_case; use tokio::sync::{mpsc, oneshot}; use tokio::task::JoinSet; @@ -400,10 +401,15 @@ mod tests { } // Returns the stream IDs in the pool, in order. - fn stream_ids(&self) -> Vec { + pub(crate) fn stream_ids(&self) -> Vec { let mut ids: Vec<_> = self.streams.lock().unwrap().iter().map(|s| s.id).collect(); ids.sort(); ids } + + // Acquire the stream lock + pub(crate) fn lock(&self) -> MutexGuard<'_, Vec> { + self.streams.lock().unwrap() + } } } From dd0334ea19996ea88aa93809ba4b0ce669bd2c3c Mon Sep 17 00:00:00 2001 From: Darren Bolduc Date: Wed, 9 Sep 2026 11:34:52 -0400 Subject: [PATCH 2/3] fix std mutex usage in test --- src/bigquery/src/write/dispatcher.rs | 31 +++++++++++++++++----------- 1 file changed, 19 insertions(+), 12 deletions(-) diff --git a/src/bigquery/src/write/dispatcher.rs b/src/bigquery/src/write/dispatcher.rs index 9da9229d44..4bdfe9e16b 100644 --- a/src/bigquery/src/write/dispatcher.rs +++ b/src/bigquery/src/write/dispatcher.rs @@ -93,7 +93,7 @@ mod tests { use crate::write::test::*; use bigquery_grpc_mock::{MockBigQueryWrite, start}; use gaxi::grpc::tonic::{Response as TonicResponse, Status as TonicStatus}; - use tokio::sync::mpsc; + use tokio::sync::{mpsc, oneshot}; use tokio::task::JoinSet; fn test_req() -> AppendRowsRequest { @@ -241,21 +241,28 @@ mod tests { let transport = Arc::new(test_transport(endpoint).await?); let pool = Arc::new(StreamPool::new(transport, 10)); let dispatcher = Arc::new(Dispatcher::new(pool.clone())); - assert_eq!(dispatcher.entry.load().id, 1); - - // Acquire the stream pool's lock to simulate a pool scaling event. - let _guard = pool.lock(); - let write = { - let d = dispatcher.clone(); - tokio::spawn(async move { d.send(test_req()).await }) - }; - - response_tx.send(Ok(convert(&test_response(1)))).await?; + // Acquire the stream pool's lock to simulate a pool scaling event. This + // needs to run in a separate thread because we don't want to hold the + // `std::sync::MutexGuard` across `await` points. + let (lock_acquired_tx, lock_acquired_rx) = oneshot::channel(); + let (release_lock_tx, release_lock_rx) = std::sync::mpsc::channel::<()>(); + std::thread::spawn(move || { + let _guard = pool.lock(); + let _ = lock_acquired_tx.send(()); + let _ = release_lock_rx.recv(); + }); + + // Wait until the lock is acquired to send a write. + lock_acquired_rx.await?; + let write = tokio::spawn(async move { dispatcher.send(test_req()).await }); // Verify the write goes through, even with the pool's lock held. + response_tx.send(Ok(convert(&test_response(1)))).await?; assert_eq!(write.await??.offset, Some(1)); - assert_eq!(dispatcher.entry.load().id, 1); + + // Release the lock + drop(release_lock_tx); Ok(()) } From b4f03c96ed9e9c52ce0d9571604487b2d4dc7734 Mon Sep 17 00:00:00 2001 From: Darren Bolduc Date: Wed, 9 Sep 2026 14:32:40 -0400 Subject: [PATCH 3/3] address review comments --- src/bigquery/src/write/dispatcher.rs | 8 ++++---- 1 file changed, 4 insertions(+), 4 deletions(-) diff --git a/src/bigquery/src/write/dispatcher.rs b/src/bigquery/src/write/dispatcher.rs index 4bdfe9e16b..7444ee4a1f 100644 --- a/src/bigquery/src/write/dispatcher.rs +++ b/src/bigquery/src/write/dispatcher.rs @@ -51,7 +51,7 @@ impl Dispatcher { /// /// Evicts and updates its cached stream on transient errors. pub(crate) async fn send(&self, req: AppendRowsRequest) -> AppendResult { - let req = req.to_proto().map_err(Error::deser)?; + let req = req.to_proto().map_err(Error::ser)?; let stream = self.entry.load_full(); let stream_id = stream.id; @@ -64,7 +64,7 @@ impl Dispatcher { let new_stream = self.pool.evict_and_replace(stream_id); // The application can `send()` multiple writes - // concurrently. Only one `send()` should update the cached + // concurrently. Only one `send()` will update the cached // stream on a transient error. let _ = self.entry.compare_and_swap(&stream, Arc::new(new_stream)); @@ -74,7 +74,7 @@ impl Dispatcher { } }?; - let resp = resp.cnv().map_err(Error::ser)?; + let resp = resp.cnv().map_err(Error::deser)?; to_result(resp) } } @@ -223,7 +223,7 @@ mod tests { } // We ran into a transient error. We should now have a new stream. Only - // one of the dispatchers should have evicted the failed stream. + // one of the callers should have evicted the failed stream. assert_eq!(dispatcher.entry.load().id, 2); assert_eq!(pool.stream_ids(), [2]);