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