use super::base::BaseWriter;
use crate::Result;
use crate::model::{
ArrowRecordBatch, ArrowSchema, BatchCommitWriteStreamsResponse, FinalizeWriteStreamResponse,
};
use crate::write::append_builder::AppendWithOffset;
use crate::write::transport::Transport;
use std::sync::Arc;
#[derive(Debug)]
pub struct PendingWriter {
pub(crate) inner: BaseWriter,
}
impl PendingWriter {
pub(crate) fn new(inner: Arc<Transport>, write_stream: String, schema: ArrowSchema) -> Self {
Self {
inner: BaseWriter::new(inner, write_stream, schema),
}
}
pub fn write_stream(&self) -> &str {
&self.inner.write_stream
}
pub fn append(&self, rows: ArrowRecordBatch) -> AppendWithOffset {
AppendWithOffset::new(
self.inner.runner.req_tx.clone(),
self.inner.append_request(rows),
)
}
pub async fn finalize(&self) -> Result<FinalizeWriteStreamResponse> {
self.inner.finalize().await
}
pub async fn commit(&self) -> Result<BatchCommitWriteStreamsResponse> {
let parent = self
.inner
.write_stream
.split_once("/streams/")
.map_or(self.inner.write_stream.as_str(), |(p, _)| p)
.to_string();
self.inner
.client
.batch_commit_write_streams()
.set_parent(parent)
.set_write_streams(vec![self.inner.write_stream.clone()])
.send()
.await
}
}
#[cfg(test)]
mod tests {
use super::super::super::runner::tests::*;
use super::super::super::transport::tests::*;
use super::*;
use bigquery_grpc_mock::{MockBigQueryWrite, start};
use gaxi::grpc::tonic::Response as TonicResponse;
use tokio::sync::mpsc;
#[tokio::test]
async fn request_fields() -> anyhow::Result<()> {
let transport = Arc::new(test_transport("http://ignored:1".to_string()).await?);
let writer = PendingWriter::new(transport, write_stream(), schema());
assert_eq!(writer.write_stream(), write_stream());
let b = writer.append(rows(1));
assert_eq!(b.req.write_stream, write_stream());
let data = b.req.arrow_rows().expect("arrow rows should be set");
let s = data.writer_schema.as_ref().expect("schema should be set");
assert_eq!(s.serialized_schema, "test");
let r = data.rows.as_ref().expect("rows should be set");
assert_eq!(r.serialized_record_batch, "1");
Ok(())
}
#[tokio::test]
async fn basic_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)));
mock.expect_finalize_write_stream()
.return_once(|_| Ok(TonicResponse::new(
bigquery_grpc_mock::google::cloud::bigquery::storage::v1::FinalizeWriteStreamResponse::default()
)));
mock.expect_batch_commit_write_streams()
.return_once(|_| Ok(TonicResponse::new(
bigquery_grpc_mock::google::cloud::bigquery::storage::v1::BatchCommitWriteStreamsResponse::default()
)));
let (endpoint, _server) = start("0.0.0.0:0", mock).await?;
let transport = Arc::new(test_transport(endpoint).await?);
let writer = PendingWriter::new(transport, write_stream(), schema());
assert_eq!(writer.write_stream(), write_stream());
response_tx.send(Ok(convert(&test_response(1)))).await?;
let resp = writer.append(rows(1)).send().await?;
assert_eq!(resp.offset, Some(1));
writer.finalize().await?;
writer.commit().await?;
Ok(())
}
fn write_stream() -> String {
"projects/p/datasets/d/tables/t/streams/s".to_string()
}
fn schema() -> ArrowSchema {
ArrowSchema::new().set_serialized_schema("test")
}
fn rows(id: i64) -> ArrowRecordBatch {
ArrowRecordBatch::new().set_serialized_record_batch(id.to_string())
}
}