1use super::DefaultWriter;
16use super::format::{Arrow, DataFormat, Proto};
17use super::generated::gapic_storage::client::BigQueryWrite;
18use super::pool::{StreamPool, StreamPoolOptions};
19use super::retry_policy::RetryOptions;
20use super::transport::Transport;
21use super::validate::{validate_stream, validate_table};
22use crate::model::write_stream::Type;
23use crate::model::{ArrowSchema, ProtoSchema, WriteStream};
24use crate::write::error::WriterBuilderError;
25use crate::write::stream_type::{ApplicationCreatedStream, DefaultStream, HasStream, Stream};
26use google_cloud_gax::backoff_policy::BackoffPolicyArg;
27use google_cloud_gax::retry_policy::RetryPolicyArg;
28use std::collections::HashMap;
29use std::marker::PhantomData;
30use std::sync::{Arc, Mutex};
31use std::time::Duration;
32
33#[derive(Clone, Debug)]
35pub struct WriterBuilder<S> {
36 pub(crate) inner: Arc<Transport>,
37 pub(crate) pools: Arc<Mutex<HashMap<String, Arc<StreamPool>>>>,
38 pub(crate) pool_options: StreamPoolOptions,
39 pub(crate) retry_options: RetryOptions,
40 op: Operation,
41 pub(crate) multiplexing: bool,
42 _stream: PhantomData<S>,
43}
44
45impl WriterBuilder<DefaultStream> {
46 pub(crate) fn new_open_default(
47 inner: Arc<Transport>,
48 pools: Arc<Mutex<HashMap<String, Arc<StreamPool>>>>,
49 pool_options: StreamPoolOptions,
50 retry_options: RetryOptions,
51 table: String,
52 ) -> Self {
53 Self {
54 inner,
55 pools,
56 pool_options,
57 retry_options,
58 op: Operation::OpenDefault { table },
59 multiplexing: false,
60 _stream: PhantomData,
61 }
62 }
63
64 pub fn with_multiplexing(mut self, enable: bool) -> Self {
88 self.multiplexing = enable;
89 self
90 }
91
92 pub fn with_retry_policy<V: Into<RetryPolicyArg>>(mut self, v: V) -> Self {
118 self.retry_options.retry_policy = v.into().into();
119 self
120 }
121
122 pub fn with_backoff_policy<V: Into<BackoffPolicyArg>>(mut self, v: V) -> Self {
147 self.retry_options.backoff_policy = v.into().into();
148 self
149 }
150
151 pub fn with_attempt_timeout(mut self, v: Duration) -> Self {
176 self.retry_options.attempt_timeout = Some(v);
177 self
178 }
179
180 pub(crate) fn make_default_writer<F: DataFormat>(
181 self,
182 write_stream: String,
183 location: String,
184 format: F,
185 ) -> DefaultWriter<F> {
186 let pool = if self.multiplexing {
187 let key = format!("{}-{}", location, format.format_name());
188 let mut pools = self.pools.lock().expect("pools lock poisoned");
189 pools
190 .entry(key)
191 .or_insert_with(|| {
192 Arc::new(StreamPool::new(
193 self.inner.clone(),
194 self.pool_options.clone(),
195 ))
196 })
197 .clone()
198 } else {
199 let options = StreamPoolOptions {
200 max_streams: 1,
201 ..Default::default()
202 };
203 Arc::new(StreamPool::new(self.inner, options))
204 };
205 DefaultWriter::new(pool, self.retry_options, write_stream, format)
206 }
207}
208
209impl<S: ApplicationCreatedStream> WriterBuilder<S> {
210 pub(crate) fn new_create(
211 inner: Arc<Transport>,
212 retry_options: RetryOptions,
213 table: String,
214 ) -> Self {
215 Self {
216 inner,
217 pools: Arc::new(Mutex::new(HashMap::new())),
218 pool_options: StreamPoolOptions::default(),
219 retry_options,
220 op: Operation::Create {
221 table,
222 stream_type: S::STREAM_TYPE,
223 },
224 multiplexing: false,
225 _stream: PhantomData,
226 }
227 }
228
229 pub(crate) fn new_attach(
230 inner: Arc<Transport>,
231 retry_options: RetryOptions,
232 write_stream: String,
233 ) -> Self {
234 Self {
235 inner,
236 pools: Arc::new(Mutex::new(HashMap::new())),
237 pool_options: StreamPoolOptions::default(),
238 retry_options,
239 op: Operation::Attach {
240 write_stream,
241 stream_type: S::STREAM_TYPE,
242 },
243 multiplexing: false,
244 _stream: PhantomData,
245 }
246 }
247}
248
249impl<S: Stream> WriterBuilder<S> {
250 pub async fn build_arrow<W>(
273 self,
274 schema: ArrowSchema,
275 ) -> std::result::Result<W, WriterBuilderError>
276 where
277 S: Stream<Writer<Arrow> = W>,
278 W: HasStream<Stream = S>,
279 {
280 self.build(Arrow { schema }).await
281 }
282
283 #[allow(dead_code)]
285 pub(crate) async fn build_proto<W>(
286 self,
287 schema: ProtoSchema,
288 ) -> std::result::Result<W, WriterBuilderError>
289 where
290 S: Stream<Writer<Proto> = W>,
291 W: HasStream<Stream = S>,
292 {
293 self.build(Proto { schema }).await
294 }
295
296 async fn build<F>(self, format: F) -> std::result::Result<S::Writer<F>, WriterBuilderError>
297 where
298 F: DataFormat,
299 {
300 let (write_stream, location) = match &self.op {
301 Operation::OpenDefault { table } => self.open_default(table).await?,
302 Operation::Create { table, stream_type } => {
303 let stream = self.create_stream(table, stream_type.clone()).await?;
304 (stream, String::new())
305 }
306 Operation::Attach {
307 write_stream,
308 stream_type,
309 } => {
310 let stream = self
311 .attach_to_stream(write_stream, stream_type.clone())
312 .await?;
313 (stream, String::new())
314 }
315 };
316 Ok(S::build(self, write_stream, location, format))
317 }
318
319 async fn open_default(
320 &self,
321 table: &str,
322 ) -> std::result::Result<(String, String), WriterBuilderError> {
323 validate_table(table)?;
324 let write_stream = format!("{table}/streams/_default");
325 let location = if self.multiplexing {
326 let client = BigQueryWrite::from_stub::<Transport>(self.inner.clone());
327 let stream = client
328 .get_write_stream()
329 .set_name(&write_stream)
330 .send()
331 .await?;
332 stream.location
333 } else {
334 String::new()
335 };
336 Ok((write_stream, location))
337 }
338
339 async fn create_stream(
340 &self,
341 table: &str,
342 stream_type: Type,
343 ) -> std::result::Result<String, WriterBuilderError> {
344 validate_table(table)?;
345
346 let client = BigQueryWrite::from_stub::<Transport>(self.inner.clone());
347 let stream = client
348 .create_write_stream()
349 .set_parent(table)
350 .set_write_stream(WriteStream::new().set_type(stream_type))
351 .send()
352 .await?;
353
354 Ok(stream.name)
355 }
356
357 async fn attach_to_stream(
358 &self,
359 write_stream: &str,
360 stream_type: Type,
361 ) -> std::result::Result<String, WriterBuilderError> {
362 validate_stream(write_stream)?;
363
364 if write_stream.ends_with("/streams/_default") {
365 return Err(WriterBuilderError::TypeMismatch {
366 expected: format!("{stream_type:?}"),
367 actual: "Default (use `open_default_stream` instead)".to_string(),
368 });
369 }
370
371 let client = BigQueryWrite::from_stub::<Transport>(self.inner.clone());
372 let stream = client
373 .get_write_stream()
374 .set_name(write_stream)
375 .send()
376 .await?;
377
378 if stream_type != stream.r#type {
379 return Err(WriterBuilderError::TypeMismatch {
380 expected: format!("{stream_type:?}"),
381 actual: format!("{:?}", stream.r#type),
382 });
383 }
384 Ok(stream.name)
385 }
386}
387
388#[derive(Clone, Debug)]
389enum Operation {
390 OpenDefault {
391 table: String,
392 },
393 Create {
394 table: String,
395 stream_type: Type,
396 },
397 Attach {
398 write_stream: String,
399 stream_type: Type,
400 },
401}
402
403#[cfg(test)]
404mod tests {
405 use super::super::format::Arrow;
406 use super::*;
407 use crate::model::write_stream::Type;
408 use crate::write::test::*;
409 use crate::write::{BufferedWriter, CommittedWriter, PendingWriter};
410 use bigquery_grpc_mock::google::cloud::bigquery::storage::v1::WriteStream as MockWriteStream;
411 use bigquery_grpc_mock::{MockBigQueryWrite, start};
412 use google_cloud_gax::retry_policy::AlwaysRetry;
413 use test_case::test_case;
414 use tokio::task::JoinHandle;
415
416 type Result<T> = std::result::Result<T, WriterBuilderError>;
417
418 #[tokio::test]
419 async fn default_stream_options() -> anyhow::Result<()> {
420 let transport = Arc::new(test_transport("http://ignored:1").await?);
421 let builder = test_open_default(transport, "projects/p/datasets/d/tables/t");
422 assert!(!builder.multiplexing);
423 assert_eq!(builder.retry_options.attempt_timeout, None);
424
425 let builder = builder
426 .with_retry_policy(AlwaysRetry)
427 .with_backoff_policy(NoBackoff)
428 .with_attempt_timeout(Duration::from_secs(10));
429 assert_eq!(
430 builder.retry_options.attempt_timeout,
431 Some(Duration::from_secs(10))
432 );
433
434 let fmt = format!("{:?}", builder.retry_options);
435 assert!(fmt.contains("AlwaysRetry"), "{fmt}");
436 assert!(fmt.contains("NoBackoff"), "{fmt}");
437
438 Ok(())
439 }
440
441 #[tokio::test]
442 async fn default() -> anyhow::Result<()> {
443 let transport = Arc::new(test_transport("http://ignored:1").await?);
444 let writer = test_open_default(transport, "projects/p/datasets/d/tables/t")
445 .build_arrow(schema())
446 .await?;
447 assert_eq!(
448 writer.write_stream,
449 "projects/p/datasets/d/tables/t/streams/_default"
450 );
451 assert_eq!(writer.format.schema, schema());
452 Ok(())
453 }
454
455 #[test_case("projects/p")]
456 #[test_case("projects/p/tables/t")]
457 #[test_case("projects/p/datasets/d/tables/")]
458 #[test_case("projects/p/instances/i/tables/t")]
459 #[test_case("projects/p/datasets/d/tables/t/streams")]
460 #[test_case("projects/p/datasets/d/tables/t/streams/_default")]
461 #[tokio::test]
462 async fn default_bad_table_format(table: &str) -> anyhow::Result<()> {
463 let transport = Arc::new(test_transport("http://ignored:1").await?);
464 let err = test_open_default(transport, table)
465 .build_arrow(schema())
466 .await
467 .expect_err("should fail locally on bad format");
468 assert!(matches!(err, WriterBuilderError::Rpc { source: e } if e.is_binding()));
469 Ok(())
470 }
471
472 async fn create_mock(stream_type: Type) -> anyhow::Result<(Arc<Transport>, JoinHandle<()>)> {
473 let mut mock = MockBigQueryWrite::new();
474 mock.expect_create_write_stream().return_once(move |req| {
475 let req = req.into_inner();
476 assert_eq!(req.parent, "projects/p/datasets/d/tables/t");
477 let ws = req.write_stream.expect("write_stream populated");
478 assert_eq!(Type::from(ws.r#type), stream_type);
479 Ok(gaxi::grpc::tonic::Response::new(MockWriteStream {
480 name: "projects/p/datasets/d/tables/t/streams/s".to_string(),
481 r#type: stream_type.value().expect("known enum value"),
482 ..Default::default()
483 }))
484 });
485 let (endpoint, server) = start("0.0.0.0:0", mock).await?;
486 let transport = Arc::new(test_transport(endpoint).await?);
487 Ok((transport, server))
488 }
489
490 #[tokio::test]
491 async fn create_committed_success() -> anyhow::Result<()> {
492 let (transport, _server) = create_mock(Type::Committed).await?;
493 let writer: CommittedWriter<Arrow> =
494 test_create(transport, "projects/p/datasets/d/tables/t")
495 .build_arrow(schema())
496 .await?;
497 assert_eq!(
498 writer.inner.write_stream,
499 "projects/p/datasets/d/tables/t/streams/s"
500 );
501 assert_eq!(writer.inner.format.schema, schema());
502 Ok(())
503 }
504
505 #[tokio::test]
506 async fn create_pending_success() -> anyhow::Result<()> {
507 let (transport, _server) = create_mock(Type::Pending).await?;
508 let writer: PendingWriter<Arrow> = test_create(transport, "projects/p/datasets/d/tables/t")
509 .build_arrow(schema())
510 .await?;
511 assert_eq!(
512 writer.inner.write_stream,
513 "projects/p/datasets/d/tables/t/streams/s"
514 );
515 assert_eq!(writer.inner.format.schema, schema());
516 Ok(())
517 }
518
519 #[tokio::test]
520 async fn create_buffered_success() -> anyhow::Result<()> {
521 let (transport, _server) = create_mock(Type::Buffered).await?;
522 let writer: BufferedWriter<Arrow> =
523 test_create(transport, "projects/p/datasets/d/tables/t")
524 .build_arrow(schema())
525 .await?;
526 assert_eq!(
527 writer.inner.write_stream,
528 "projects/p/datasets/d/tables/t/streams/s"
529 );
530 assert_eq!(writer.inner.format.schema, schema());
531 Ok(())
532 }
533
534 #[test_case("projects/p")]
535 #[test_case("projects/p/tables/t")]
536 #[test_case("projects/p/datasets/d/tables/")]
537 #[test_case("projects/p/instances/i/tables/t")]
538 #[test_case("projects/p/datasets/d/tables/t/streams")]
539 #[test_case("projects/p/datasets/d/tables/t/streams/_default")]
540 #[tokio::test]
541 async fn create_bad_table_format(table: &str) -> anyhow::Result<()> {
542 let transport = Arc::new(test_transport("http://ignored:1").await?);
543 let res: Result<PendingWriter<Arrow>> =
544 test_create(transport, table).build_arrow(schema()).await;
545 let err = res.expect_err("should fail locally on bad format");
546 assert!(matches!(err, WriterBuilderError::Rpc { source: e } if e.is_binding()));
547 Ok(())
548 }
549
550 async fn attach_mock(stream_type: Type) -> anyhow::Result<(Arc<Transport>, JoinHandle<()>)> {
551 let mut mock = MockBigQueryWrite::new();
552 mock.expect_get_write_stream().return_once(move |req| {
553 let req = req.into_inner();
554 assert_eq!(req.name, "projects/p/datasets/d/tables/t/streams/s");
555 Ok(gaxi::grpc::tonic::Response::new(MockWriteStream {
556 name: "projects/p/datasets/d/tables/t/streams/s".to_string(),
557 r#type: stream_type.value().expect("known enum value"),
558 ..Default::default()
559 }))
560 });
561 let (endpoint, server) = start("0.0.0.0:0", mock).await?;
562 let transport = Arc::new(test_transport(endpoint).await?);
563 Ok((transport, server))
564 }
565
566 #[tokio::test]
567 async fn attach_committed_success() -> anyhow::Result<()> {
568 let (transport, _server) = attach_mock(Type::Committed).await?;
569 let writer: CommittedWriter<Arrow> =
570 test_attach(transport, "projects/p/datasets/d/tables/t/streams/s")
571 .build_arrow(schema())
572 .await?;
573 assert_eq!(
574 writer.inner.write_stream,
575 "projects/p/datasets/d/tables/t/streams/s"
576 );
577 assert_eq!(writer.inner.format.schema, schema());
578 Ok(())
579 }
580
581 #[tokio::test]
582 async fn attach_pending_success() -> anyhow::Result<()> {
583 let (transport, _server) = attach_mock(Type::Pending).await?;
584 let writer: PendingWriter<Arrow> =
585 test_attach(transport, "projects/p/datasets/d/tables/t/streams/s")
586 .build_arrow(schema())
587 .await?;
588 assert_eq!(
589 writer.inner.write_stream,
590 "projects/p/datasets/d/tables/t/streams/s"
591 );
592 assert_eq!(writer.inner.format.schema, schema());
593 Ok(())
594 }
595
596 #[tokio::test]
597 async fn attach_buffered_success() -> anyhow::Result<()> {
598 let (transport, _server) = attach_mock(Type::Buffered).await?;
599 let writer: BufferedWriter<Arrow> =
600 test_attach(transport, "projects/p/datasets/d/tables/t/streams/s")
601 .build_arrow(schema())
602 .await?;
603 assert_eq!(
604 writer.inner.write_stream,
605 "projects/p/datasets/d/tables/t/streams/s"
606 );
607 assert_eq!(writer.inner.format.schema, schema());
608 Ok(())
609 }
610
611 #[test_case("projects/p")]
612 #[test_case("projects/p/tables/t")]
613 #[test_case("projects/p/datasets/d/tables/t")]
614 #[test_case("projects/p/datasets/d/tables/t/streams/")]
615 #[tokio::test]
616 async fn attach_bad_stream_format(stream: &str) -> anyhow::Result<()> {
617 let transport = Arc::new(test_transport("http://ignored:1").await?);
618 let res: Result<CommittedWriter<Arrow>> =
619 test_attach(transport, stream).build_arrow(schema()).await;
620 let err = res.expect_err("should fail locally on bad format");
621 assert!(matches!(err, WriterBuilderError::Rpc { source: e } if e.is_binding()));
622 Ok(())
623 }
624
625 #[tokio::test]
626 async fn attach_stream_type_mismatch() -> anyhow::Result<()> {
627 let (transport, _server) = attach_mock(Type::Buffered).await?;
628 let res: Result<CommittedWriter<Arrow>> =
629 test_attach(transport, "projects/p/datasets/d/tables/t/streams/s")
630 .build_arrow(schema())
631 .await;
632 let err = res.expect_err("should return type mismatch error");
633 assert!(matches!(err, WriterBuilderError::TypeMismatch { .. }));
634 assert!(err.to_string().contains("stream type mismatch: requested"));
635 Ok(())
636 }
637
638 #[tokio::test]
639 async fn attach_default_stream_rejected() -> anyhow::Result<()> {
640 let transport = Arc::new(test_transport("http://ignored:1").await?);
641 let default_stream = "projects/p/datasets/d/tables/t/streams/_default";
642
643 let res: Result<CommittedWriter<Arrow>> = test_attach(transport.clone(), default_stream)
644 .build_arrow(schema())
645 .await;
646 let err = res.expect_err("should reject attaching CommittedWriter to _default");
647 assert!(
648 matches!(
649 &err,
650 WriterBuilderError::TypeMismatch { expected, actual }
651 if expected == "Committed" && actual.contains("Default")
652 ),
653 "unexpected error: {err:?}"
654 );
655 assert!(
656 err.to_string().contains(
657 "stream type mismatch: requested Committed, but matched resource yields Default"
658 ),
659 "unexpected display: {err}"
660 );
661
662 let res: Result<PendingWriter<Arrow>> = test_attach(transport.clone(), default_stream)
663 .build_arrow(schema())
664 .await;
665 let err = res.expect_err("should reject attaching PendingWriter to _default");
666 assert!(
667 matches!(
668 &err,
669 WriterBuilderError::TypeMismatch { expected, actual }
670 if expected == "Pending" && actual.contains("Default")
671 ),
672 "unexpected error: {err:?}"
673 );
674
675 let res: Result<BufferedWriter<Arrow>> = test_attach(transport, default_stream)
676 .build_arrow(schema())
677 .await;
678 let err = res.expect_err("should reject attaching BufferedWriter to _default");
679 assert!(
680 matches!(
681 &err,
682 WriterBuilderError::TypeMismatch { expected, actual }
683 if expected == "Buffered" && actual.contains("Default")
684 ),
685 "unexpected error: {err:?}"
686 );
687
688 Ok(())
689 }
690
691 fn test_open_default(transport: Arc<Transport>, table: &str) -> WriterBuilder<DefaultStream> {
692 let pools = Arc::new(Mutex::new(HashMap::new()));
693 WriterBuilder::new_open_default(
694 transport,
695 pools,
696 StreamPoolOptions::default(),
697 test_retry_options(),
698 table.to_string(),
699 )
700 }
701
702 fn test_create<S: ApplicationCreatedStream>(
703 transport: Arc<Transport>,
704 table: &str,
705 ) -> WriterBuilder<S> {
706 WriterBuilder::new_create(transport, test_retry_options(), table.to_string())
707 }
708
709 fn test_attach<S: ApplicationCreatedStream>(
710 transport: Arc<Transport>,
711 write_stream: &str,
712 ) -> WriterBuilder<S> {
713 WriterBuilder::new_attach(transport, test_retry_options(), write_stream.to_string())
714 }
715}