Skip to main content

risingwave_connector/sink/
big_query.rs

1// Copyright 2023 RisingWave Labs
2//
3// Licensed under the Apache License, Version 2.0 (the "License");
4// you may not use this file except in compliance with the License.
5// You may obtain a copy of the License at
6//
7//     http://www.apache.org/licenses/LICENSE-2.0
8//
9// Unless required by applicable law or agreed to in writing, software
10// distributed under the License is distributed on an "AS IS" BASIS,
11// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12// See the License for the specific language governing permissions and
13// limitations under the License.
14
15use core::pin::Pin;
16use core::time::Duration;
17use std::collections::{BTreeMap, HashMap, VecDeque};
18
19use anyhow::{Context, anyhow};
20use async_trait::async_trait;
21use base64::Engine;
22use base64::prelude::BASE64_STANDARD;
23use futures::future::pending;
24use futures::prelude::Future;
25use futures::{Stream, StreamExt};
26use futures_async_stream::try_stream;
27use gcp_bigquery_client::Client;
28use gcp_bigquery_client::error::BQError;
29use gcp_bigquery_client::model::query_request::QueryRequest;
30use gcp_bigquery_client::model::query_response::ResultSet;
31use gcp_bigquery_client::model::table::Table;
32use gcp_bigquery_client::model::table_field_schema::TableFieldSchema;
33use gcp_bigquery_client::model::table_schema::TableSchema;
34use google_cloud_bigquery::grpc::apiv1::conn_pool::ConnectionManager;
35use google_cloud_gax::conn::{ConnectionOptions, Environment};
36use google_cloud_gax::grpc::{Request, Response, Status};
37use google_cloud_googleapis::cloud::bigquery::storage::v1::append_rows_request::{
38    MissingValueInterpretation, ProtoData, Rows as AppendRowsRequestRows,
39};
40use google_cloud_googleapis::cloud::bigquery::storage::v1::{
41    AppendRowsRequest, AppendRowsResponse, ProtoRows, ProtoSchema,
42};
43use google_cloud_pubsub::client::google_cloud_auth;
44use google_cloud_pubsub::client::google_cloud_auth::credentials::CredentialsFile;
45use phf::{Set, phf_set};
46use prost_reflect::{FieldDescriptor, MessageDescriptor};
47use prost_types::{
48    DescriptorProto, FieldDescriptorProto, FileDescriptorProto, FileDescriptorSet,
49    field_descriptor_proto,
50};
51use risingwave_common::array::{Op, StreamChunk};
52use risingwave_common::catalog::{Field, Schema};
53use risingwave_common::types::DataType;
54use serde::Deserialize;
55use serde_with::{DisplayFromStr, serde_as};
56use simd_json::prelude::ArrayTrait;
57use tokio::sync::mpsc;
58use url::Url;
59use uuid::Uuid;
60use with_options::WithOptions;
61use yup_oauth2::ServiceAccountKey;
62
63use super::encoder::{ProtoEncoder, ProtoHeader, RowEncoder, SerTo};
64use super::log_store::{LogStoreReadItem, TruncateOffset};
65use super::{
66    LogSinker, SINK_TYPE_APPEND_ONLY, SINK_TYPE_OPTION, SINK_TYPE_UPSERT, SinkError, SinkLogReader,
67};
68use crate::aws_utils::load_file_descriptor_from_s3;
69use crate::connector_common::AwsAuthProps;
70use crate::enforce_secret::EnforceSecret;
71use crate::sink::{Result, Sink, SinkParam, SinkWriterParam};
72
73pub const BIGQUERY_SINK: &str = "bigquery";
74pub const CHANGE_TYPE: &str = "_CHANGE_TYPE";
75const DEFAULT_GRPC_CHANNEL_NUMS: usize = 4;
76const CONNECT_TIMEOUT: Option<Duration> = Some(Duration::from_secs(30));
77const CONNECTION_TIMEOUT: Option<Duration> = None;
78const BIGQUERY_SEND_FUTURE_BUFFER_MAX_SIZE: usize = 65536;
79// < 10MB, we set 8MB
80const MAX_ROW_SIZE: usize = 8 * 1024 * 1024;
81
82#[serde_as]
83#[derive(Deserialize, Debug, Clone, WithOptions)]
84pub struct BigQueryCommon {
85    #[serde(rename = "bigquery.local.path")]
86    pub local_path: Option<String>,
87    #[serde(rename = "bigquery.s3.path")]
88    pub s3_path: Option<String>,
89    #[serde(rename = "bigquery.project")]
90    pub project: String,
91    #[serde(rename = "bigquery.dataset")]
92    pub dataset: String,
93    #[serde(rename = "bigquery.table")]
94    pub table: String,
95    #[serde(default, alias = "create_table_if_not_exists")] // default false
96    #[serde_as(as = "DisplayFromStr")]
97    pub auto_create: bool,
98    #[serde(rename = "bigquery.credentials")]
99    pub credentials: Option<String>,
100}
101
102impl EnforceSecret for BigQueryCommon {
103    const ENFORCE_SECRET_PROPERTIES: Set<&'static str> = phf_set! {
104        "bigquery.credentials",
105    };
106}
107
108struct BigQueryFutureManager {
109    // `offset_queue` holds the Some corresponding to each future.
110    // When TruncateOffset is barrier, the num is 0, we don't need to wait for the return of `resp_stream`.
111    // When TruncateOffset is chunk:
112    // 1. chunk has no rows. we didn't send, the num is 0, we don't need to wait for the return of `resp_stream`.
113    // 2. chunk is less than `MAX_ROW_SIZE`, we only sent once, the num is 1 and we only have to wait once for `resp_stream`.
114    // 3. chunk is less than `MAX_ROW_SIZE`, we only sent n, the num is n and we need to wait n times for r.
115    offset_queue: VecDeque<(TruncateOffset, usize)>,
116    resp_stream: Pin<Box<dyn Stream<Item = Result<()>> + Send>>,
117}
118impl BigQueryFutureManager {
119    pub fn new(
120        max_future_num: usize,
121        resp_stream: impl Stream<Item = Result<()>> + Send + 'static,
122    ) -> Self {
123        let offset_queue = VecDeque::with_capacity(max_future_num);
124        Self {
125            offset_queue,
126            resp_stream: Box::pin(resp_stream),
127        }
128    }
129
130    pub fn add_offset(&mut self, offset: TruncateOffset, resp_num: usize) {
131        self.offset_queue.push_back((offset, resp_num));
132    }
133
134    pub async fn next_offset(&mut self) -> Result<TruncateOffset> {
135        if let Some((_offset, remaining_resp_num)) = self.offset_queue.front_mut() {
136            if *remaining_resp_num == 0 {
137                return Ok(self.offset_queue.pop_front().unwrap().0);
138            }
139            while *remaining_resp_num > 0 {
140                self.resp_stream
141                    .next()
142                    .await
143                    .ok_or_else(|| SinkError::BigQuery(anyhow::anyhow!("end of stream")))??;
144                *remaining_resp_num -= 1;
145            }
146            Ok(self.offset_queue.pop_front().unwrap().0)
147        } else {
148            pending().await
149        }
150    }
151}
152pub struct BigQueryLogSinker {
153    writer: BigQuerySinkWriter,
154    bigquery_future_manager: BigQueryFutureManager,
155    future_num: usize,
156}
157impl BigQueryLogSinker {
158    pub fn new(
159        writer: BigQuerySinkWriter,
160        resp_stream: impl Stream<Item = Result<()>> + Send + 'static,
161        future_num: usize,
162    ) -> Self {
163        Self {
164            writer,
165            bigquery_future_manager: BigQueryFutureManager::new(future_num, resp_stream),
166            future_num,
167        }
168    }
169}
170
171#[async_trait]
172impl LogSinker for BigQueryLogSinker {
173    async fn consume_log_and_sink(mut self, mut log_reader: impl SinkLogReader) -> Result<!> {
174        log_reader.start_from(None).await?;
175        loop {
176            tokio::select!(
177                offset = self.bigquery_future_manager.next_offset() => {
178                        log_reader.truncate(offset?)?;
179                }
180                item_result = log_reader.next_item(), if self.bigquery_future_manager.offset_queue.len() <= self.future_num => {
181                    let (epoch, item) = item_result?;
182                    match item {
183                        LogStoreReadItem::StreamChunk { chunk_id, chunk } => {
184                            let resp_num = self.writer.write_chunk(chunk)?;
185                            self.bigquery_future_manager
186                                .add_offset(TruncateOffset::Chunk { epoch, chunk_id },resp_num);
187                        }
188                        LogStoreReadItem::Barrier { .. } => {
189                            self.bigquery_future_manager
190                                .add_offset(TruncateOffset::Barrier { epoch },0);
191                        }
192                    }
193                }
194            )
195        }
196    }
197}
198
199impl BigQueryCommon {
200    async fn build_client(&self, aws_auth_props: &AwsAuthProps) -> Result<Client> {
201        let auth_json = self.get_auth_json_from_path(aws_auth_props).await?;
202
203        let service_account =
204            if let Ok(auth_json_from_base64) = BASE64_STANDARD.decode(auth_json.clone()) {
205                serde_json::from_slice::<ServiceAccountKey>(&auth_json_from_base64)
206            } else {
207                serde_json::from_str::<ServiceAccountKey>(&auth_json)
208            }
209            .map_err(|e| SinkError::BigQuery(e.into()))?;
210
211        let client: Client = Client::from_service_account_key(service_account, false)
212            .await
213            .map_err(|err| SinkError::BigQuery(anyhow::anyhow!(err)))?;
214        Ok(client)
215    }
216
217    async fn build_writer_client(
218        &self,
219        aws_auth_props: &AwsAuthProps,
220    ) -> Result<(StorageWriterClient, impl Stream<Item = Result<()>> + use<>)> {
221        let auth_json = self.get_auth_json_from_path(aws_auth_props).await?;
222
223        let credentials_file =
224            if let Ok(auth_json_from_base64) = BASE64_STANDARD.decode(auth_json.clone()) {
225                serde_json::from_slice::<CredentialsFile>(&auth_json_from_base64)
226            } else {
227                serde_json::from_str::<CredentialsFile>(&auth_json)
228            }
229            .map_err(|e| SinkError::BigQuery(e.into()))?;
230
231        StorageWriterClient::new(credentials_file).await
232    }
233
234    async fn get_auth_json_from_path(&self, aws_auth_props: &AwsAuthProps) -> Result<String> {
235        if let Some(credentials) = &self.credentials {
236            Ok(credentials.clone())
237        } else if let Some(local_path) = &self.local_path {
238            std::fs::read_to_string(local_path)
239                .map_err(|err| SinkError::BigQuery(anyhow::anyhow!(err)))
240        } else if let Some(s3_path) = &self.s3_path {
241            let url =
242                Url::parse(s3_path).map_err(|err| SinkError::BigQuery(anyhow::anyhow!(err)))?;
243            let auth_vec = load_file_descriptor_from_s3(&url, aws_auth_props)
244                .await
245                .map_err(|err| SinkError::BigQuery(anyhow::anyhow!(err)))?;
246            Ok(String::from_utf8(auth_vec).map_err(|e| SinkError::BigQuery(e.into()))?)
247        } else {
248            Err(SinkError::BigQuery(anyhow::anyhow!(
249                "`bigquery.local.path` and `bigquery.s3.path` set at least one, configure as needed."
250            )))
251        }
252    }
253}
254
255#[serde_as]
256#[derive(Clone, Debug, Deserialize, WithOptions)]
257pub struct BigQueryConfig {
258    #[serde(flatten)]
259    pub common: BigQueryCommon,
260    #[serde(flatten)]
261    pub aws_auth_props: AwsAuthProps,
262    pub r#type: String, // accept "append-only" or "upsert"
263
264    #[serde(flatten)]
265    pub unknown_fields: std::collections::HashMap<String, String>,
266}
267
268crate::impl_sink_unknown_fields!(BigQueryConfig);
269
270impl EnforceSecret for BigQueryConfig {
271    fn enforce_one(prop: &str) -> crate::error::ConnectorResult<()> {
272        BigQueryCommon::enforce_one(prop)?;
273        AwsAuthProps::enforce_one(prop)?;
274        Ok(())
275    }
276}
277
278impl BigQueryConfig {
279    pub fn from_btreemap(properties: BTreeMap<String, String>) -> Result<Self> {
280        let config =
281            serde_json::from_value::<BigQueryConfig>(serde_json::to_value(properties).unwrap())
282                .map_err(|e| SinkError::Config(anyhow!(e)))?;
283        if config.r#type != SINK_TYPE_APPEND_ONLY && config.r#type != SINK_TYPE_UPSERT {
284            return Err(SinkError::Config(anyhow!(
285                "`{}` must be {}, or {}",
286                SINK_TYPE_OPTION,
287                SINK_TYPE_APPEND_ONLY,
288                SINK_TYPE_UPSERT
289            )));
290        }
291        Ok(config)
292    }
293}
294
295#[derive(Debug)]
296pub struct BigQuerySink {
297    pub config: BigQueryConfig,
298    schema: Schema,
299    pk_indices: Vec<usize>,
300    is_append_only: bool,
301}
302
303impl EnforceSecret for BigQuerySink {
304    fn enforce_secret<'a>(
305        prop_iter: impl Iterator<Item = &'a str>,
306    ) -> crate::error::ConnectorResult<()> {
307        for prop in prop_iter {
308            BigQueryConfig::enforce_one(prop)?;
309        }
310        Ok(())
311    }
312}
313
314impl BigQuerySink {
315    pub fn new(
316        config: BigQueryConfig,
317        schema: Schema,
318        pk_indices: Vec<usize>,
319        is_append_only: bool,
320    ) -> Result<Self> {
321        Ok(Self {
322            config,
323            schema,
324            pk_indices,
325            is_append_only,
326        })
327    }
328}
329
330impl BigQuerySink {
331    fn is_decimal_type_compatible(bigquery_type: &str) -> bool {
332        // BigQuery INFORMATION_SCHEMA reports parameterized decimal columns as
333        // `NUMERIC(p, s)` or `BIGNUMERIC(p, s)`. RisingWave `Decimal` does not
334        // carry typmod, so schema validation accepts the BigQuery decimal type
335        // family instead of requiring exact string equality.
336        let normalized = bigquery_type.trim().to_ascii_uppercase();
337        matches!(
338            normalized
339                .split_once('(')
340                .map_or(normalized.as_str(), |(prefix, _)| prefix),
341            "NUMERIC" | "BIGNUMERIC"
342        )
343    }
344
345    fn is_data_type_compatible(rw_data_type: &DataType, bigquery_type: &str) -> Result<bool> {
346        if matches!(rw_data_type, DataType::Decimal) {
347            return Ok(Self::is_decimal_type_compatible(bigquery_type));
348        }
349
350        Ok(Self::get_string_and_check_support_from_datatype(rw_data_type)? == bigquery_type)
351    }
352
353    fn check_column_name_and_type(
354        &self,
355        big_query_columns_desc: HashMap<String, String>,
356    ) -> Result<()> {
357        let rw_fields_name = self.schema.fields();
358        if big_query_columns_desc.is_empty() {
359            return Err(SinkError::BigQuery(anyhow::anyhow!(
360                "Cannot find table in bigquery"
361            )));
362        }
363        if rw_fields_name.len().ne(&big_query_columns_desc.len()) {
364            return Err(SinkError::BigQuery(anyhow::anyhow!(
365                "The length of the RisingWave column {} must be equal to the length of the bigquery column {}",
366                rw_fields_name.len(),
367                big_query_columns_desc.len()
368            )));
369        }
370
371        for i in rw_fields_name {
372            let value = big_query_columns_desc.get(&i.name).ok_or_else(|| {
373                SinkError::BigQuery(anyhow::anyhow!(
374                    "Column `{:?}` on RisingWave side is not found on BigQuery side.",
375                    i.name
376                ))
377            })?;
378            let data_type_string = Self::get_string_and_check_support_from_datatype(&i.data_type)?;
379            if !Self::is_data_type_compatible(&i.data_type, value)? {
380                return Err(SinkError::BigQuery(anyhow::anyhow!(
381                    "Data type mismatch for column `{:?}`. BigQuery side: `{:?}`, RisingWave side: `{:?}`. ",
382                    i.name,
383                    value,
384                    data_type_string
385                )));
386            };
387        }
388        Ok(())
389    }
390
391    fn get_string_and_check_support_from_datatype(rw_data_type: &DataType) -> Result<String> {
392        match rw_data_type {
393            DataType::Boolean => Ok("BOOL".to_owned()),
394            DataType::Int16 => Ok("INT64".to_owned()),
395            DataType::Int32 => Ok("INT64".to_owned()),
396            DataType::Int64 => Ok("INT64".to_owned()),
397            DataType::Float32 => Err(SinkError::BigQuery(anyhow::anyhow!(
398                "REAL is not supported for BigQuery sink. Please convert to FLOAT64 or other supported types."
399            ))),
400            DataType::Float64 => Ok("FLOAT64".to_owned()),
401            DataType::Decimal => Ok("NUMERIC".to_owned()),
402            DataType::Date => Ok("DATE".to_owned()),
403            DataType::Varchar => Ok("STRING".to_owned()),
404            DataType::Time => Ok("TIME".to_owned()),
405            DataType::Timestamp => Ok("DATETIME".to_owned()),
406            DataType::Timestamptz => Ok("TIMESTAMP".to_owned()),
407            DataType::Interval => Ok("INTERVAL".to_owned()),
408            DataType::Struct(structs) => {
409                let mut elements_vec = vec![];
410                for (name, datatype) in structs.iter() {
411                    let element_string =
412                        Self::get_string_and_check_support_from_datatype(datatype)?;
413                    elements_vec.push(format!("{} {}", name, element_string));
414                }
415                Ok(format!("STRUCT<{}>", elements_vec.join(", ")))
416            }
417            DataType::List(l) => {
418                let element_string = Self::get_string_and_check_support_from_datatype(l.elem())?;
419                Ok(format!("ARRAY<{}>", element_string))
420            }
421            DataType::Bytea => Ok("BYTES".to_owned()),
422            DataType::Jsonb => Ok("JSON".to_owned()),
423            DataType::Variant => Err(SinkError::BigQuery(anyhow::anyhow!(
424                "VARIANT is not supported for BigQuery sink."
425            ))),
426            DataType::Serial => Ok("INT64".to_owned()),
427            DataType::Int256 => Err(SinkError::BigQuery(anyhow::anyhow!(
428                "INT256 is not supported for BigQuery sink."
429            ))),
430            DataType::Map(_) => Err(SinkError::BigQuery(anyhow::anyhow!(
431                "MAP is not supported for BigQuery sink."
432            ))),
433            DataType::Vector(_) => Err(SinkError::BigQuery(anyhow::anyhow!(
434                "VECTOR is not supported for BigQuery sink."
435            ))),
436        }
437    }
438
439    fn map_field(rw_field: &Field) -> Result<TableFieldSchema> {
440        let tfs = match &rw_field.data_type {
441            DataType::Boolean => TableFieldSchema::bool(&rw_field.name),
442            DataType::Int16 | DataType::Int32 | DataType::Int64 | DataType::Serial => {
443                TableFieldSchema::integer(&rw_field.name)
444            }
445            DataType::Float32 => {
446                return Err(SinkError::BigQuery(anyhow::anyhow!(
447                    "REAL is not supported for BigQuery sink. Please convert to FLOAT64 or other supported types."
448                )));
449            }
450            DataType::Float64 => TableFieldSchema::float(&rw_field.name),
451            DataType::Decimal => TableFieldSchema::numeric(&rw_field.name),
452            DataType::Date => TableFieldSchema::date(&rw_field.name),
453            DataType::Varchar => TableFieldSchema::string(&rw_field.name),
454            DataType::Time => TableFieldSchema::time(&rw_field.name),
455            DataType::Timestamp => TableFieldSchema::date_time(&rw_field.name),
456            DataType::Timestamptz => TableFieldSchema::timestamp(&rw_field.name),
457            DataType::Interval => {
458                return Err(SinkError::BigQuery(anyhow::anyhow!(
459                    "INTERVAL is not supported for BigQuery sink. Please convert to VARCHAR or other supported types."
460                )));
461            }
462            DataType::Struct(st) => {
463                let mut sub_fields = Vec::with_capacity(st.len());
464                for (name, dt) in st.iter() {
465                    let rw_field = Field::with_name(dt.clone(), name);
466                    let field = Self::map_field(&rw_field)?;
467                    sub_fields.push(field);
468                }
469                TableFieldSchema::record(&rw_field.name, sub_fields)
470            }
471            DataType::List(lt) => {
472                let inner_field =
473                    Self::map_field(&Field::with_name(lt.elem().clone(), &rw_field.name))?;
474                TableFieldSchema {
475                    mode: Some("REPEATED".to_owned()),
476                    ..inner_field
477                }
478            }
479
480            DataType::Bytea => TableFieldSchema::bytes(&rw_field.name),
481            DataType::Jsonb => TableFieldSchema::json(&rw_field.name),
482            DataType::Variant => {
483                return Err(SinkError::BigQuery(anyhow::anyhow!(
484                    "VARIANT is not supported for BigQuery sink."
485                )));
486            }
487            DataType::Int256 => {
488                return Err(SinkError::BigQuery(anyhow::anyhow!(
489                    "INT256 is not supported for BigQuery sink."
490                )));
491            }
492            DataType::Map(_) => {
493                return Err(SinkError::BigQuery(anyhow::anyhow!(
494                    "MAP is not supported for BigQuery sink."
495                )));
496            }
497            DataType::Vector(_) => {
498                return Err(SinkError::BigQuery(anyhow::anyhow!(
499                    "VECTOR is not supported for BigQuery sink."
500                )));
501            }
502        };
503        Ok(tfs)
504    }
505
506    async fn create_table(
507        &self,
508        client: &Client,
509        project_id: &str,
510        dataset_id: &str,
511        table_id: &str,
512        fields: &Vec<Field>,
513    ) -> Result<Table> {
514        let dataset = client
515            .dataset()
516            .get(project_id, dataset_id)
517            .await
518            .map_err(|e| SinkError::BigQuery(e.into()))?;
519        let fields: Vec<_> = fields.iter().map(Self::map_field).collect::<Result<_>>()?;
520        let table = Table::from_dataset(&dataset, table_id, TableSchema::new(fields));
521
522        client
523            .table()
524            .create(table)
525            .await
526            .map_err(|e| SinkError::BigQuery(e.into()))
527    }
528}
529
530impl Sink for BigQuerySink {
531    type LogSinker = BigQueryLogSinker;
532
533    const SINK_NAME: &'static str = BIGQUERY_SINK;
534
535    crate::impl_validate_sink_unknown_fields!();
536
537    async fn new_log_sinker(&self, _writer_param: SinkWriterParam) -> Result<Self::LogSinker> {
538        let (writer, resp_stream) = BigQuerySinkWriter::new(
539            self.config.clone(),
540            self.schema.clone(),
541            self.pk_indices.clone(),
542            self.is_append_only,
543        )
544        .await?;
545        Ok(BigQueryLogSinker::new(
546            writer,
547            resp_stream,
548            BIGQUERY_SEND_FUTURE_BUFFER_MAX_SIZE,
549        ))
550    }
551
552    async fn validate(&self) -> Result<()> {
553        risingwave_common::license::Feature::BigQuerySink
554            .check_available()
555            .map_err(|e| anyhow::anyhow!(e))?;
556        if !self.is_append_only && self.pk_indices.is_empty() {
557            return Err(SinkError::Config(anyhow!(
558                "Primary key not defined for upsert bigquery sink (please define in `primary_key` field)"
559            )));
560        }
561        let client = self
562            .config
563            .common
564            .build_client(&self.config.aws_auth_props)
565            .await?;
566        let BigQueryCommon {
567            project: project_id,
568            dataset: dataset_id,
569            table: table_id,
570            ..
571        } = &self.config.common;
572
573        if self.config.common.auto_create {
574            match client
575                .table()
576                .get(project_id, dataset_id, table_id, None)
577                .await
578            {
579                Err(BQError::ResponseError { error }) if error.error.code == 404 => {
580                    // early return: no need to query schema to check column and type
581                    return self
582                        .create_table(
583                            &client,
584                            project_id,
585                            dataset_id,
586                            table_id,
587                            &self.schema.fields,
588                        )
589                        .await
590                        .map(|_| ());
591                }
592                Err(e) => return Err(SinkError::BigQuery(e.into())),
593                _ => {}
594            }
595        }
596
597        let rs = client
598            .job()
599            .query(
600                &self.config.common.project,
601                QueryRequest::new(format!(
602                    "SELECT column_name, data_type FROM `{}.{}.INFORMATION_SCHEMA.COLUMNS` WHERE table_name = '{}'",
603                    project_id, dataset_id, table_id,
604                )),
605            ).await.map_err(|e| SinkError::BigQuery(e.into()))?;
606        let mut rs = ResultSet::new_from_query_response(rs);
607
608        let mut big_query_schema = HashMap::default();
609        while rs.next_row() {
610            big_query_schema.insert(
611                rs.get_string_by_name("column_name")
612                    .map_err(|e| SinkError::BigQuery(e.into()))?
613                    .ok_or_else(|| {
614                        SinkError::BigQuery(anyhow::anyhow!("Cannot find column_name"))
615                    })?,
616                rs.get_string_by_name("data_type")
617                    .map_err(|e| SinkError::BigQuery(e.into()))?
618                    .ok_or_else(|| {
619                        SinkError::BigQuery(anyhow::anyhow!("Cannot find column_name"))
620                    })?,
621            );
622        }
623
624        self.check_column_name_and_type(big_query_schema)?;
625        Ok(())
626    }
627}
628
629pub struct BigQuerySinkWriter {
630    pub config: BigQueryConfig,
631    #[expect(dead_code)]
632    schema: Schema,
633    #[expect(dead_code)]
634    pk_indices: Vec<usize>,
635    client: StorageWriterClient,
636    is_append_only: bool,
637    row_encoder: ProtoEncoder,
638    writer_pb_schema: ProtoSchema,
639    #[expect(dead_code)]
640    message_descriptor: MessageDescriptor,
641    write_stream: String,
642    proto_field: Option<FieldDescriptor>,
643}
644
645impl TryFrom<SinkParam> for BigQuerySink {
646    type Error = SinkError;
647
648    fn try_from(param: SinkParam) -> std::result::Result<Self, Self::Error> {
649        let schema = param.schema();
650        let pk_indices = param.downstream_pk_or_empty();
651        let config = BigQueryConfig::from_btreemap(param.properties)?;
652        BigQuerySink::new(config, schema, pk_indices, param.sink_type.is_append_only())
653    }
654}
655
656impl BigQuerySinkWriter {
657    pub async fn new(
658        config: BigQueryConfig,
659        schema: Schema,
660        pk_indices: Vec<usize>,
661        is_append_only: bool,
662    ) -> Result<(Self, impl Stream<Item = Result<()>>)> {
663        let (client, resp_stream) = config
664            .common
665            .build_writer_client(&config.aws_auth_props)
666            .await?;
667        let mut descriptor_proto = build_protobuf_schema(
668            schema
669                .fields()
670                .iter()
671                .map(|f| (f.name.as_str(), &f.data_type)),
672            config.common.table.clone(),
673        )?;
674
675        if !is_append_only {
676            let field = FieldDescriptorProto {
677                name: Some(CHANGE_TYPE.to_owned()),
678                number: Some((schema.len() + 1) as i32),
679                r#type: Some(field_descriptor_proto::Type::String.into()),
680                ..Default::default()
681            };
682            descriptor_proto.field.push(field);
683        }
684
685        let descriptor_pool = build_protobuf_descriptor_pool(&descriptor_proto)?;
686        let message_descriptor = descriptor_pool
687            .get_message_by_name(&config.common.table)
688            .ok_or_else(|| {
689                SinkError::BigQuery(anyhow::anyhow!(
690                    "Can't find message proto {}",
691                    config.common.table
692                ))
693            })?;
694        let proto_field = if !is_append_only {
695            let proto_field = message_descriptor
696                .get_field_by_name(CHANGE_TYPE)
697                .ok_or_else(|| {
698                    SinkError::BigQuery(anyhow::anyhow!("Can't find {}", CHANGE_TYPE))
699                })?;
700            Some(proto_field)
701        } else {
702            None
703        };
704        let row_encoder = ProtoEncoder::new(
705            schema.clone(),
706            None,
707            message_descriptor.clone(),
708            ProtoHeader::None,
709        )?;
710        Ok((
711            Self {
712                write_stream: format!(
713                    "projects/{}/datasets/{}/tables/{}/streams/_default",
714                    config.common.project, config.common.dataset, config.common.table
715                ),
716                config,
717                schema,
718                pk_indices,
719                client,
720                is_append_only,
721                row_encoder,
722                message_descriptor,
723                proto_field,
724                writer_pb_schema: ProtoSchema {
725                    proto_descriptor: Some(descriptor_proto.clone()),
726                },
727            },
728            resp_stream,
729        ))
730    }
731
732    fn append_only(&mut self, chunk: StreamChunk) -> Result<Vec<Vec<u8>>> {
733        let mut serialized_rows: Vec<Vec<u8>> = Vec::with_capacity(chunk.capacity());
734        for (op, row) in chunk.rows() {
735            if op != Op::Insert {
736                continue;
737            }
738            serialized_rows.push(self.row_encoder.encode(row)?.ser_to()?)
739        }
740        Ok(serialized_rows)
741    }
742
743    fn upsert(&mut self, chunk: StreamChunk) -> Result<Vec<Vec<u8>>> {
744        let mut serialized_rows: Vec<Vec<u8>> = Vec::with_capacity(chunk.capacity());
745        for (op, row) in chunk.rows() {
746            if op == Op::UpdateDelete {
747                continue;
748            }
749            let mut pb_row = self.row_encoder.encode(row)?;
750            match op {
751                Op::Insert => pb_row
752                    .message
753                    .try_set_field(
754                        self.proto_field.as_ref().unwrap(),
755                        prost_reflect::Value::String("UPSERT".to_owned()),
756                    )
757                    .map_err(|e| SinkError::BigQuery(e.into()))?,
758                Op::Delete => pb_row
759                    .message
760                    .try_set_field(
761                        self.proto_field.as_ref().unwrap(),
762                        prost_reflect::Value::String("DELETE".to_owned()),
763                    )
764                    .map_err(|e| SinkError::BigQuery(e.into()))?,
765                Op::UpdateDelete => continue,
766                Op::UpdateInsert => pb_row
767                    .message
768                    .try_set_field(
769                        self.proto_field.as_ref().unwrap(),
770                        prost_reflect::Value::String("UPSERT".to_owned()),
771                    )
772                    .map_err(|e| SinkError::BigQuery(e.into()))?,
773            };
774
775            serialized_rows.push(pb_row.ser_to()?)
776        }
777        Ok(serialized_rows)
778    }
779
780    fn write_chunk(&mut self, chunk: StreamChunk) -> Result<usize> {
781        let serialized_rows = if self.is_append_only {
782            self.append_only(chunk)?
783        } else {
784            self.upsert(chunk)?
785        };
786        if serialized_rows.is_empty() {
787            return Ok(0);
788        }
789        let mut result = Vec::new();
790        let mut result_inner = Vec::new();
791        let mut size_count = 0;
792        for i in serialized_rows {
793            size_count += i.len();
794            if size_count > MAX_ROW_SIZE {
795                result.push(result_inner);
796                result_inner = Vec::new();
797                size_count = i.len();
798            }
799            result_inner.push(i);
800        }
801        if !result_inner.is_empty() {
802            result.push(result_inner);
803        }
804        let len = result.len();
805        for serialized_rows in result {
806            let rows = AppendRowsRequestRows::ProtoRows(ProtoData {
807                writer_schema: Some(self.writer_pb_schema.clone()),
808                rows: Some(ProtoRows { serialized_rows }),
809            });
810            self.client.append_rows(rows, self.write_stream.clone())?;
811        }
812        Ok(len)
813    }
814}
815
816#[try_stream(ok = (), error = SinkError)]
817pub async fn resp_to_stream(
818    resp_stream: impl Future<
819        Output = std::result::Result<
820            Response<google_cloud_gax::grpc::Streaming<AppendRowsResponse>>,
821            Status,
822        >,
823    >
824    + 'static
825    + Send,
826) {
827    let mut resp_stream = resp_stream
828        .await
829        .map_err(|e| SinkError::BigQuery(e.into()))?
830        .into_inner();
831    loop {
832        match resp_stream
833            .message()
834            .await
835            .map_err(|e| SinkError::BigQuery(e.into()))?
836        {
837            Some(append_rows_response) => {
838                if !append_rows_response.row_errors.is_empty() {
839                    return Err(SinkError::BigQuery(anyhow::anyhow!(
840                        "bigquery insert error {:?}",
841                        append_rows_response.row_errors
842                    )));
843                }
844                if let Some(google_cloud_googleapis::cloud::bigquery::storage::v1::append_rows_response::Response::Error(status)) = append_rows_response.response{
845                            return Err(SinkError::BigQuery(anyhow::anyhow!(
846                                "bigquery insert error {:?}",
847                                status
848                            )));
849                        }
850                yield ();
851            }
852            None => {
853                return Err(SinkError::BigQuery(anyhow::anyhow!(
854                    "bigquery insert error: end of resp stream",
855                )));
856            }
857        }
858    }
859}
860
861struct StorageWriterClient {
862    #[expect(dead_code)]
863    environment: Environment,
864    request_sender: mpsc::UnboundedSender<AppendRowsRequest>,
865}
866impl StorageWriterClient {
867    pub async fn new(
868        credentials: CredentialsFile,
869    ) -> Result<(Self, impl Stream<Item = Result<()>>)> {
870        let ts_grpc = google_cloud_auth::token::DefaultTokenSourceProvider::new_with_credentials(
871            Self::bigquery_grpc_auth_config(),
872            Box::new(credentials),
873        )
874        .await
875        .map_err(|e| SinkError::BigQuery(e.into()))?;
876        let conn_options = ConnectionOptions {
877            connect_timeout: CONNECT_TIMEOUT,
878            timeout: CONNECTION_TIMEOUT,
879            ..Default::default()
880        };
881        let environment = Environment::GoogleCloud(Box::new(ts_grpc));
882        let conn = ConnectionManager::new(DEFAULT_GRPC_CHANNEL_NUMS, &environment, &conn_options)
883            .await
884            .map_err(|e| SinkError::BigQuery(e.into()))?;
885        let mut client = conn.writer();
886
887        let (tx, rx) = mpsc::unbounded_channel();
888        let stream = tokio_stream::wrappers::UnboundedReceiverStream::new(rx);
889
890        let resp = async move { client.append_rows(Request::new(stream)).await };
891        let resp_stream = resp_to_stream(resp);
892
893        Ok((
894            StorageWriterClient {
895                environment,
896                request_sender: tx,
897            },
898            resp_stream,
899        ))
900    }
901
902    pub fn append_rows(&mut self, row: AppendRowsRequestRows, write_stream: String) -> Result<()> {
903        let append_req = AppendRowsRequest {
904            write_stream,
905            offset: None,
906            trace_id: Uuid::new_v4().hyphenated().to_string(),
907            missing_value_interpretations: HashMap::default(),
908            rows: Some(row),
909            default_missing_value_interpretation: MissingValueInterpretation::DefaultValue as i32,
910        };
911        self.request_sender
912            .send(append_req)
913            .map_err(|e| SinkError::BigQuery(e.into()))?;
914        Ok(())
915    }
916
917    fn bigquery_grpc_auth_config() -> google_cloud_auth::project::Config<'static> {
918        let mut auth_config = google_cloud_auth::project::Config::default();
919        auth_config =
920            auth_config.with_audience(google_cloud_bigquery::grpc::apiv1::conn_pool::AUDIENCE);
921        auth_config =
922            auth_config.with_scopes(&google_cloud_bigquery::grpc::apiv1::conn_pool::SCOPES);
923        auth_config
924    }
925}
926
927fn build_protobuf_descriptor_pool(desc: &DescriptorProto) -> Result<prost_reflect::DescriptorPool> {
928    let file_descriptor = FileDescriptorProto {
929        message_type: vec![desc.clone()],
930        name: Some("bigquery".to_owned()),
931        ..Default::default()
932    };
933
934    prost_reflect::DescriptorPool::from_file_descriptor_set(FileDescriptorSet {
935        file: vec![file_descriptor],
936    })
937    .context("failed to build descriptor pool")
938    .map_err(SinkError::BigQuery)
939}
940
941fn build_protobuf_schema<'a>(
942    fields: impl Iterator<Item = (&'a str, &'a DataType)>,
943    name: String,
944) -> Result<DescriptorProto> {
945    let mut proto = DescriptorProto {
946        name: Some(name),
947        ..Default::default()
948    };
949    let mut struct_vec = vec![];
950    let field_vec = fields
951        .enumerate()
952        .map(|(index, (name, data_type))| {
953            let (field, des_proto) =
954                build_protobuf_field(data_type, (index + 1) as i32, name.to_owned())?;
955            if let Some(sv) = des_proto {
956                struct_vec.push(sv);
957            }
958            Ok(field)
959        })
960        .collect::<Result<Vec<_>>>()?;
961    proto.field = field_vec;
962    proto.nested_type = struct_vec;
963    Ok(proto)
964}
965
966fn build_protobuf_field(
967    data_type: &DataType,
968    index: i32,
969    name: String,
970) -> Result<(FieldDescriptorProto, Option<DescriptorProto>)> {
971    let mut field = FieldDescriptorProto {
972        name: Some(name.clone()),
973        number: Some(index),
974        ..Default::default()
975    };
976    match data_type {
977        DataType::Boolean => field.r#type = Some(field_descriptor_proto::Type::Bool.into()),
978        DataType::Int32 => field.r#type = Some(field_descriptor_proto::Type::Int32.into()),
979        DataType::Int16 | DataType::Int64 => {
980            field.r#type = Some(field_descriptor_proto::Type::Int64.into())
981        }
982        DataType::Float64 => field.r#type = Some(field_descriptor_proto::Type::Double.into()),
983        DataType::Decimal => field.r#type = Some(field_descriptor_proto::Type::String.into()),
984        DataType::Date => field.r#type = Some(field_descriptor_proto::Type::Int32.into()),
985        DataType::Varchar => field.r#type = Some(field_descriptor_proto::Type::String.into()),
986        DataType::Time => field.r#type = Some(field_descriptor_proto::Type::String.into()),
987        DataType::Timestamp => field.r#type = Some(field_descriptor_proto::Type::String.into()),
988        DataType::Timestamptz => field.r#type = Some(field_descriptor_proto::Type::String.into()),
989        DataType::Interval => field.r#type = Some(field_descriptor_proto::Type::String.into()),
990        DataType::Struct(s) => {
991            field.r#type = Some(field_descriptor_proto::Type::Message.into());
992            let name = format!("Struct{}", name);
993            let sub_proto = build_protobuf_schema(s.iter(), name.clone())?;
994            field.type_name = Some(name);
995            return Ok((field, Some(sub_proto)));
996        }
997        DataType::List(l) => {
998            let (mut field, proto) = build_protobuf_field(l.elem(), index, name)?;
999            field.label = Some(field_descriptor_proto::Label::Repeated.into());
1000            return Ok((field, proto));
1001        }
1002        DataType::Bytea => field.r#type = Some(field_descriptor_proto::Type::Bytes.into()),
1003        DataType::Jsonb => field.r#type = Some(field_descriptor_proto::Type::String.into()),
1004        DataType::Variant => {
1005            return Err(SinkError::BigQuery(anyhow::anyhow!("Don't support Variant")));
1006        }
1007        DataType::Serial => field.r#type = Some(field_descriptor_proto::Type::Int64.into()),
1008        DataType::Float32 | DataType::Int256 => {
1009            return Err(SinkError::BigQuery(anyhow::anyhow!(
1010                "Don't support Float32 and Int256"
1011            )));
1012        }
1013        DataType::Map(_) => return Err(SinkError::BigQuery(anyhow::anyhow!("Don't support Map"))),
1014        DataType::Vector(_) => {
1015            return Err(SinkError::BigQuery(anyhow::anyhow!("Don't support Vector")));
1016        }
1017    }
1018    Ok((field, None))
1019}
1020
1021#[cfg(test)]
1022mod test {
1023
1024    use std::assert_matches;
1025    use std::collections::HashMap;
1026
1027    use risingwave_common::catalog::{Field, Schema};
1028    use risingwave_common::types::{DataType, StructType};
1029
1030    use crate::connector_common::AwsAuthProps;
1031    use crate::sink::big_query::{
1032        BigQueryCommon, BigQueryConfig, BigQuerySink, build_protobuf_descriptor_pool,
1033        build_protobuf_schema,
1034    };
1035
1036    #[tokio::test]
1037    async fn test_type_check() {
1038        let big_query_type_string = "ARRAY<STRUCT<v1 ARRAY<INT64>, v2 STRUCT<v1 INT64, v2 INT64>>>";
1039        let rw_datatype = DataType::list(DataType::Struct(StructType::new(vec![
1040            ("v1".to_owned(), DataType::Int64.list()),
1041            (
1042                "v2".to_owned(),
1043                DataType::Struct(StructType::new(vec![
1044                    ("v1".to_owned(), DataType::Int64),
1045                    ("v2".to_owned(), DataType::Int64),
1046                ])),
1047            ),
1048        ])));
1049        assert_eq!(
1050            BigQuerySink::get_string_and_check_support_from_datatype(&rw_datatype).unwrap(),
1051            big_query_type_string
1052        );
1053    }
1054
1055    #[tokio::test]
1056    async fn test_schema_check() {
1057        let schema = Schema {
1058            fields: vec![
1059                Field::with_name(DataType::Int64, "v1"),
1060                Field::with_name(DataType::Float64, "v2"),
1061                Field::with_name(
1062                    DataType::list(DataType::Struct(StructType::new(vec![
1063                        ("v1".to_owned(), DataType::Int64.list()),
1064                        (
1065                            "v3".to_owned(),
1066                            DataType::Struct(StructType::new(vec![
1067                                ("v1".to_owned(), DataType::Int64),
1068                                ("v2".to_owned(), DataType::Int64),
1069                            ])),
1070                        ),
1071                    ]))),
1072                    "v3",
1073                ),
1074            ],
1075        };
1076        let fields = schema
1077            .fields()
1078            .iter()
1079            .map(|f| (f.name.as_str(), &f.data_type));
1080        let desc = build_protobuf_schema(fields, "t1".to_owned()).unwrap();
1081        let pool = build_protobuf_descriptor_pool(&desc).unwrap();
1082        let t1_message = pool.get_message_by_name("t1").unwrap();
1083        assert_matches!(
1084            t1_message.get_field_by_name("v1").unwrap().kind(),
1085            prost_reflect::Kind::Int64
1086        );
1087        assert_matches!(
1088            t1_message.get_field_by_name("v2").unwrap().kind(),
1089            prost_reflect::Kind::Double
1090        );
1091        assert_matches!(
1092            t1_message.get_field_by_name("v3").unwrap().kind(),
1093            prost_reflect::Kind::Message(_)
1094        );
1095
1096        let v3_message = pool.get_message_by_name("t1.Structv3").unwrap();
1097        assert_matches!(
1098            v3_message.get_field_by_name("v1").unwrap().kind(),
1099            prost_reflect::Kind::Int64
1100        );
1101        assert!(v3_message.get_field_by_name("v1").unwrap().is_list());
1102
1103        let v3_v3_message = pool.get_message_by_name("t1.Structv3.Structv3").unwrap();
1104        assert_matches!(
1105            v3_v3_message.get_field_by_name("v1").unwrap().kind(),
1106            prost_reflect::Kind::Int64
1107        );
1108        assert_matches!(
1109            v3_v3_message.get_field_by_name("v2").unwrap().kind(),
1110            prost_reflect::Kind::Int64
1111        );
1112    }
1113
1114    #[test]
1115    fn test_decimal_type_family_compatibility() {
1116        assert!(BigQuerySink::is_decimal_type_compatible("NUMERIC"));
1117        assert!(BigQuerySink::is_decimal_type_compatible("numeric(31, 2)"));
1118        assert!(BigQuerySink::is_decimal_type_compatible("BIGNUMERIC"));
1119        assert!(BigQuerySink::is_decimal_type_compatible(
1120            "bignumeric(35, 12)"
1121        ));
1122        assert!(!BigQuerySink::is_decimal_type_compatible("STRING"));
1123    }
1124
1125    #[test]
1126    fn test_decimal_schema_check_accepts_parameterized_numeric_types() {
1127        let sink = BigQuerySink {
1128            config: BigQueryConfig {
1129                common: BigQueryCommon {
1130                    local_path: None,
1131                    s3_path: None,
1132                    project: "project".to_owned(),
1133                    dataset: "dataset".to_owned(),
1134                    table: "table".to_owned(),
1135                    auto_create: false,
1136                    credentials: None,
1137                },
1138                aws_auth_props: AwsAuthProps {
1139                    region: None,
1140                    endpoint: None,
1141                    access_key: None,
1142                    secret_key: None,
1143                    session_token: None,
1144                    arn: None,
1145                    external_id: None,
1146                    profile: None,
1147                    msk_signer_timeout_sec: None,
1148                },
1149                r#type: "append-only".to_owned(),
1150                unknown_fields: Default::default(),
1151            },
1152            schema: Schema {
1153                fields: vec![Field::with_name(DataType::Decimal, "capitalizedcost")],
1154            },
1155            pk_indices: vec![],
1156            is_append_only: true,
1157        };
1158
1159        sink.check_column_name_and_type(HashMap::from([(
1160            "capitalizedcost".to_owned(),
1161            "NUMERIC(31, 2)".to_owned(),
1162        )]))
1163        .unwrap();
1164
1165        sink.check_column_name_and_type(HashMap::from([(
1166            "capitalizedcost".to_owned(),
1167            "BIGNUMERIC(35, 12)".to_owned(),
1168        )]))
1169        .unwrap();
1170    }
1171}