use super::append_future::AppendFuture;
use super::append_response::to_result;
use super::error::AppendError;
use super::runner::WriteRequest;
use crate::Error;
use crate::model::AppendRowsRequest;
use gaxi::prost::{FromProto, ToProto};
use tokio::sync::{mpsc, oneshot};
#[derive(Clone, Debug)]
pub struct AppendWithOffset {
req_tx: mpsc::UnboundedSender<WriteRequest>,
pub(crate) req: AppendRowsRequest,
}
impl AppendWithOffset {
#[allow(dead_code)]
pub(crate) fn new(req_tx: mpsc::UnboundedSender<WriteRequest>, req: AppendRowsRequest) -> Self {
Self { req_tx, req }
}
pub fn set_offset(mut self, offset: i64) -> Self {
self.req.offset = Some(offset);
self
}
pub fn send(self) -> AppendFuture {
let (tx, rx) = oneshot::channel();
let (resp_tx, resp_rx) = oneshot::channel();
let req = match self.req.to_proto().map_err(Error::deser) {
Ok(req) => req,
Err(e) => {
let _ = tx.send(Err(e.into()));
return AppendFuture::new(rx);
}
};
let write = WriteRequest { req, resp_tx };
let _ = self.req_tx.send(write);
tokio::spawn(async move {
let res = async {
let resp = resp_rx
.await
.map_err(|_| AppendError::UnexpectedEndOfStream)??;
let resp = resp.cnv().map_err(Error::ser)?;
to_result(resp)
}
.await;
let _ = tx.send(res);
});
AppendFuture::new(rx)
}
}
#[derive(Clone, Debug)]
pub struct Append {
req_tx: mpsc::UnboundedSender<WriteRequest>,
pub(crate) req: AppendRowsRequest,
}
impl Append {
pub(crate) fn new(req_tx: mpsc::UnboundedSender<WriteRequest>, req: AppendRowsRequest) -> Self {
Self { req_tx, req }
}
pub fn send(self) -> AppendFuture {
let (tx, rx) = oneshot::channel();
tokio::spawn(async move {
let (resp_tx, resp_rx) = oneshot::channel();
let res = async move {
let req = self.req.to_proto().map_err(Error::deser)?;
let write = WriteRequest { req, resp_tx };
let _ = self.req_tx.send(write);
let resp = resp_rx
.await
.map_err(|_| AppendError::UnexpectedEndOfStream)??;
let resp = resp.cnv().map_err(Error::ser)?;
to_result(resp)
}
.await;
let _ = tx.send(res);
});
AppendFuture::new(rx)
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::google::cloud::bigquery::storage::v1;
use crate::google::cloud::bigquery::storage::v1::append_rows_response::{
AppendResult, Response,
};
use crate::model::TableSchema;
#[tokio::test]
async fn success() -> anyhow::Result<()> {
let (req_tx, mut req_rx) = mpsc::unbounded_channel();
let req = AppendRowsRequest::new().set_write_stream(write_stream());
let builder = Append::new(req_tx, req);
let handle = tokio::spawn(async move { builder.send().await });
let write = req_rx.recv().await.expect("should receive request");
assert_eq!(write.req.write_stream, write_stream());
let resp = v1::AppendRowsResponse {
response: Some(Response::AppendResult(AppendResult::default())),
write_stream: write_stream(),
updated_schema: Some(v1::TableSchema::default()),
..Default::default()
};
write
.resp_tx
.send(Ok(resp))
.expect("sending on channel always succeeds");
let resp = handle.await??;
assert_eq!(resp.offset, None);
assert_eq!(resp.updated_schema, Some(TableSchema::default()));
Ok(())
}
#[tokio::test]
async fn stream_closed() -> anyhow::Result<()> {
let (req_tx, req_rx) = mpsc::unbounded_channel();
let req = AppendRowsRequest::new().set_write_stream(write_stream());
let builder = Append::new(req_tx, req);
let handle = tokio::spawn(async move { builder.send().await });
drop(req_rx);
let err = handle.await?.expect_err("should return an error");
assert!(matches!(err, AppendError::UnexpectedEndOfStream));
Ok(())
}
#[tokio::test]
async fn rpc_error() -> anyhow::Result<()> {
let (req_tx, mut req_rx) = mpsc::unbounded_channel();
let req = AppendRowsRequest::new().set_write_stream(write_stream());
let builder = Append::new(req_tx, req);
let handle = tokio::spawn(async move { builder.send().await });
let write = req_rx.recv().await.expect("should receive request");
let append_err: AppendError = Error::io("fail").into();
write
.resp_tx
.send(Err(append_err))
.expect("sending on channel always succeeds");
let err = handle.await?.expect_err("should return an error");
assert!(matches!(err, AppendError::Rpc { source: _ }));
Ok(())
}
#[tokio::test]
async fn row_errors() -> anyhow::Result<()> {
let (req_tx, mut req_rx) = mpsc::unbounded_channel();
let req = AppendRowsRequest::new().set_write_stream(write_stream());
let builder = Append::new(req_tx, req);
let handle = tokio::spawn(async move { builder.send().await });
let write = req_rx.recv().await.expect("should receive request");
let row_error = v1::RowError {
index: 42,
code: v1::row_error::RowErrorCode::FieldsError as i32,
message: "fail".to_string(),
};
let resp = v1::AppendRowsResponse {
row_errors: vec![row_error],
write_stream: write_stream(),
..Default::default()
};
write
.resp_tx
.send(Ok(resp))
.expect("sending on channel always succeeds");
let err = handle.await?.expect_err("should return an error");
assert!(matches!(err, AppendError::RowErrors(_)));
Ok(())
}
#[tokio::test]
async fn offset_success() -> anyhow::Result<()> {
let (req_tx, mut req_rx) = mpsc::unbounded_channel();
let req = AppendRowsRequest::new().set_write_stream(write_stream());
let builder = AppendWithOffset::new(req_tx, req).set_offset(100);
let future = builder.send();
let write = req_rx.recv().await.expect("should receive request");
assert_eq!(write.req.offset, Some(100));
let resp = v1::AppendRowsResponse {
response: Some(Response::AppendResult(AppendResult::default())),
write_stream: write_stream(),
..Default::default()
};
write
.resp_tx
.send(Ok(resp))
.expect("sending on channel always succeeds");
let resp = future.await?;
assert_eq!(resp.offset, None);
assert_eq!(resp.updated_schema, None);
Ok(())
}
#[tokio::test]
async fn offset_stream_closed() -> anyhow::Result<()> {
let (req_tx, req_rx) = mpsc::unbounded_channel();
let req = AppendRowsRequest::new().set_write_stream(write_stream());
let builder = AppendWithOffset::new(req_tx, req).set_offset(100);
let future = builder.send();
drop(req_rx);
let err = future.await.expect_err("should return an error");
assert!(matches!(err, AppendError::UnexpectedEndOfStream));
Ok(())
}
#[tokio::test]
async fn offset_rpc_error() -> anyhow::Result<()> {
let (req_tx, mut req_rx) = mpsc::unbounded_channel();
let req = AppendRowsRequest::new().set_write_stream(write_stream());
let builder = AppendWithOffset::new(req_tx, req).set_offset(100);
let future = builder.send();
let write = req_rx.recv().await.expect("should receive request");
let append_err: AppendError = Error::io("fail").into();
write
.resp_tx
.send(Err(append_err))
.expect("sending on channel always succeeds");
let err = future.await.expect_err("should return an error");
assert!(matches!(err, AppendError::Rpc { source: _ }));
Ok(())
}
#[tokio::test]
async fn offset_row_errors() -> anyhow::Result<()> {
let (req_tx, mut req_rx) = mpsc::unbounded_channel();
let req = AppendRowsRequest::new().set_write_stream(write_stream());
let builder = AppendWithOffset::new(req_tx, req).set_offset(100);
let future = builder.send();
let write = req_rx.recv().await.expect("should receive request");
let row_error = v1::RowError {
index: 42,
code: v1::row_error::RowErrorCode::FieldsError as i32,
message: "fail".to_string(),
};
let resp = v1::AppendRowsResponse {
row_errors: vec![row_error],
write_stream: write_stream(),
..Default::default()
};
write
.resp_tx
.send(Ok(resp))
.expect("sending on channel always succeeds");
let err = future.await.expect_err("should return an error");
assert!(matches!(err, AppendError::RowErrors(_)));
Ok(())
}
#[tokio::test(flavor = "multi_thread", worker_threads = 8)]
async fn synchronous_queueing() -> anyhow::Result<()> {
const NUM_WRITES: i64 = 1000;
let (req_tx, mut req_rx) = mpsc::unbounded_channel();
let write_handle = tokio::spawn(async move {
let mut writes = tokio::task::JoinSet::new();
for i in 0..NUM_WRITES {
writes.spawn(
AppendWithOffset::new(req_tx.clone(), AppendRowsRequest::new())
.set_offset(i)
.send(),
);
}
let _ = writes.join_all().await;
});
for i in 0..NUM_WRITES {
let write = req_rx.recv().await.expect("should receive request");
assert_eq!(write.req.offset, Some(i), "received out of order write");
}
write_handle.await?;
Ok(())
}
fn write_stream() -> String {
"projects/p/datasets/d/tables/t/streams/_default".to_string()
}
}