use super::error::{AppendError, AppendResult};
use super::stream::Stream;
use super::transport::Transport;
use crate::Result;
use crate::google::cloud::bigquery::storage::v1::{AppendRowsRequest, AppendRowsResponse};
use gaxi::grpc::from_status::to_gax_error;
use gaxi::grpc::tonic::{Status as TonicStatus, Streaming};
use std::collections::VecDeque;
use std::sync::Arc;
use tokio::sync::{mpsc, oneshot};
use tokio::task::JoinHandle;
type TonicResult<T> = std::result::Result<T, TonicStatus>;
#[derive(Debug)]
pub(crate) struct WriteRequest {
pub(crate) req: AppendRowsRequest,
pub(crate) resp_tx: oneshot::Sender<AppendResult<AppendRowsResponse>>,
}
#[derive(Debug)]
pub(crate) struct Runner {
pub(crate) req_tx: mpsc::UnboundedSender<WriteRequest>,
#[allow(dead_code)]
pub(crate) handle: JoinHandle<()>,
}
impl Runner {
pub(crate) fn new(inner: Arc<Transport>) -> Self {
let (req_tx, req_rx) = mpsc::unbounded_channel();
let handle = tokio::spawn(async move {
run_stream_task(inner, req_rx).await;
});
Runner { req_tx, handle }
}
}
async fn run_stream_task(inner: Arc<Transport>, mut req_rx: mpsc::UnboundedReceiver<WriteRequest>) {
let Some(initial_req) = req_rx.recv().await else {
return;
};
let mut resp_txs = VecDeque::new();
resp_txs.push_back(initial_req.resp_tx);
let Stream {
mut stream,
request_tx,
} = match Stream::new(inner, initial_req.req).await {
Ok(s) => s,
Err(e) => {
process_gax_response(&mut resp_txs, Err(e));
return;
}
};
loop {
tokio::select! {
req = req_rx.recv() => {
match req {
Some(r) => {
resp_txs.push_back(r.resp_tx);
let _ = request_tx.send(r.req).await;
}
None => break drain_stream(stream, resp_txs).await,
}
}
resp = stream.message() => {
match resp.transpose() {
Some(r) => process_response(&mut resp_txs, r),
None => break,
}
}
}
}
}
async fn drain_stream(
mut stream: Streaming<AppendRowsResponse>,
mut resp_txs: VecDeque<oneshot::Sender<AppendResult<AppendRowsResponse>>>,
) {
while let Some(r) = stream.message().await.transpose() {
process_response(&mut resp_txs, r);
}
}
fn process_response(
resp_txs: &mut VecDeque<oneshot::Sender<AppendResult<AppendRowsResponse>>>,
resp: TonicResult<AppendRowsResponse>,
) {
process_gax_response(resp_txs, resp.map_err(to_gax_error))
}
fn process_gax_response(
resp_txs: &mut VecDeque<oneshot::Sender<AppendResult<AppendRowsResponse>>>,
resp: Result<AppendRowsResponse>,
) {
let resp_tx = resp_txs
.pop_front()
.expect("the service sends one response per request");
let _ = resp_tx.send(resp.map_err(AppendError::from));
}
#[cfg(test)]
pub(crate) mod tests {
use super::super::transport::tests::*;
use super::*;
use crate::google::cloud::bigquery::storage::v1::append_rows_response::{
AppendResult, Response,
};
use bigquery_grpc_mock::{MockBigQueryWrite, start};
use gaxi::grpc::tonic::Response as TonicResponse;
use google_cloud_gax::error::rpc::Code;
#[tokio::test]
async fn no_requests() -> anyhow::Result<()> {
let (_, response_rx) = mpsc::channel(1);
let mut mock = MockBigQueryWrite::new();
mock.expect_append_rows()
.return_once(|_| 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 Runner { req_tx, handle } = Runner::new(transport);
drop(req_tx);
handle.await?;
Ok(())
}
#[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(|_| 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 Runner { req_tx, handle } = Runner::new(transport);
let (resp_tx1, resp_rx1) = oneshot::channel();
let write1 = WriteRequest {
req: test_request(1),
resp_tx: resp_tx1,
};
req_tx.send(write1)?;
let (resp_tx2, resp_rx2) = oneshot::channel();
let write2 = WriteRequest {
req: test_request(2),
resp_tx: resp_tx2,
};
req_tx.send(write2)?;
response_tx.send(Ok(convert(&test_response(1)))).await?;
let resp1 = resp_rx1.await??;
assert_eq!(resp1, test_response(1));
let (resp_tx3, resp_rx3) = oneshot::channel();
let write3 = WriteRequest {
req: test_request(3),
resp_tx: resp_tx3,
};
req_tx.send(write3)?;
response_tx.send(Ok(convert(&test_response(2)))).await?;
let resp2 = resp_rx2.await??;
assert_eq!(resp2, test_response(2));
response_tx.send(Ok(convert(&test_response(3)))).await?;
let resp3 = resp_rx3.await??;
assert_eq!(resp3, test_response(3));
drop(req_tx);
drop(response_tx);
handle.await?;
Ok(())
}
#[tokio::test]
async fn error_starting_stream() -> anyhow::Result<()> {
let mut mock = MockBigQueryWrite::new();
mock.expect_append_rows()
.return_once(|_| Err(TonicStatus::failed_precondition("fail")));
let (endpoint, _server) = start("0.0.0.0:0", mock).await?;
let transport = Arc::new(test_transport(endpoint).await?);
let Runner { req_tx, handle } = Runner::new(transport);
let (resp_tx, resp_rx) = oneshot::channel();
let write = WriteRequest {
req: test_request(1),
resp_tx,
};
req_tx.send(write)?;
let resp = resp_rx.await?;
let Err(AppendError::Rpc { source: err }) = resp else {
anyhow::bail!("expected an RPC error, got: {resp:?}");
};
let Some(status) = err.status() else {
anyhow::bail!("expected a status, got: {err:?}");
};
assert_eq!(status.code, Code::FailedPrecondition);
assert_eq!(status.message, "fail");
drop(req_tx);
handle.await?;
Ok(())
}
#[tokio::test]
async fn error_mid_stream() -> anyhow::Result<()> {
let (response_tx, response_rx) = mpsc::channel(10);
let mut mock = MockBigQueryWrite::new();
mock.expect_append_rows()
.return_once(|_| 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 Runner { req_tx, handle } = Runner::new(transport);
let (resp_tx1, resp_rx1) = oneshot::channel();
let write1 = WriteRequest {
req: test_request(1),
resp_tx: resp_tx1,
};
req_tx.send(write1)?;
let (resp_tx2, resp_rx2) = oneshot::channel();
let write2 = WriteRequest {
req: test_request(2),
resp_tx: resp_tx2,
};
req_tx.send(write2)?;
let (resp_tx3, resp_rx3) = oneshot::channel();
let write3 = WriteRequest {
req: test_request(3),
resp_tx: resp_tx3,
};
req_tx.send(write3)?;
response_tx.send(Ok(convert(&test_response(1)))).await?;
let resp1 = resp_rx1.await??;
assert_eq!(resp1, test_response(1));
response_tx
.send(Err(TonicStatus::failed_precondition("fail")))
.await?;
let resp2 = resp_rx2.await?;
let Err(AppendError::Rpc { source: err }) = resp2 else {
anyhow::bail!("expected an RPC error, got: {resp2:?}");
};
let Some(status) = err.status() else {
anyhow::bail!("expected a status, got: {err:?}");
};
assert_eq!(status.code, Code::FailedPrecondition);
assert_eq!(status.message, "fail");
let _resp3 = resp_rx3.await.expect_err("channel should be closed");
drop(req_tx);
drop(response_tx);
handle.await?;
Ok(())
}
#[tokio::test]
async fn sender_dropped_mid_stream() -> anyhow::Result<()> {
let (response_tx, response_rx) = mpsc::channel(10);
let mut mock = MockBigQueryWrite::new();
mock.expect_append_rows()
.return_once(|_| 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 Runner { req_tx, handle } = Runner::new(transport);
let (resp_tx1, resp_rx1) = oneshot::channel();
let write1 = WriteRequest {
req: test_request(1),
resp_tx: resp_tx1,
};
req_tx.send(write1)?;
let (resp_tx2, resp_rx2) = oneshot::channel();
let write2 = WriteRequest {
req: test_request(2),
resp_tx: resp_tx2,
};
req_tx.send(write2)?;
let (resp_tx3, resp_rx3) = oneshot::channel();
let write3 = WriteRequest {
req: test_request(3),
resp_tx: resp_tx3,
};
req_tx.send(write3)?;
response_tx.send(Ok(convert(&test_response(1)))).await?;
let resp1 = resp_rx1.await??;
assert_eq!(resp1, test_response(1));
drop(req_tx);
response_tx.send(Ok(convert(&test_response(2)))).await?;
let resp2 = resp_rx2.await??;
assert_eq!(resp2, test_response(2));
response_tx.send(Ok(convert(&test_response(3)))).await?;
let resp3 = resp_rx3.await??;
assert_eq!(resp3, test_response(3));
drop(response_tx);
handle.await?;
Ok(())
}
#[tokio::test]
async fn unexpected_end_of_stream() -> anyhow::Result<()> {
let (response_tx, response_rx) = mpsc::channel(10);
let mut mock = MockBigQueryWrite::new();
mock.expect_append_rows()
.return_once(|_| 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 Runner { req_tx, handle } = Runner::new(transport);
let (resp_tx1, resp_rx1) = oneshot::channel();
let write1 = WriteRequest {
req: test_request(1),
resp_tx: resp_tx1,
};
req_tx.send(write1)?;
let (resp_tx2, resp_rx2) = oneshot::channel();
let write2 = WriteRequest {
req: test_request(2),
resp_tx: resp_tx2,
};
req_tx.send(write2)?;
let (resp_tx3, resp_rx3) = oneshot::channel();
let write3 = WriteRequest {
req: test_request(3),
resp_tx: resp_tx3,
};
req_tx.send(write3)?;
response_tx.send(Ok(convert(&test_response(1)))).await?;
let resp1 = resp_rx1.await??;
assert_eq!(resp1, test_response(1));
drop(response_tx);
let _resp2 = resp_rx2.await.expect_err("channel should be closed");
let _resp3 = resp_rx3.await.expect_err("channel should be closed");
handle.await?;
Ok(())
}
pub(crate) fn test_request(index: i64) -> AppendRowsRequest {
AppendRowsRequest {
write_stream: "projects/p/datasets/d/tables/t/streams/s".to_string(),
offset: Some(index),
..Default::default()
}
}
pub(crate) fn test_response(index: i64) -> AppendRowsResponse {
AppendRowsResponse {
response: Some(Response::AppendResult(AppendResult {
offset: Some(index),
})),
write_stream: "projects/p/datasets/d/tables/t/streams/s".to_string(),
..Default::default()
}
}
}