1use std::collections::{BTreeMap, HashMap, HashSet};
16
17use anyhow::{Context, anyhow};
18use async_trait::async_trait;
19use base64::Engine;
20use base64::engine::general_purpose;
21use bytes::{BufMut, Bytes, BytesMut};
22use risingwave_common::array::{Op, StreamChunk};
23use risingwave_common::catalog::Schema;
24use risingwave_common::types::DataType;
25use serde::{Deserialize, Serialize};
26use serde_json::Value;
27use serde_with::{DisplayFromStr, serde_as};
28use thiserror_ext::AsReport;
29use with_options::WithOptions;
30
31use super::doris_starrocks_connector::{
32 DORIS_DELETE_SIGN, DORIS_SUCCESS_STATUS, HeaderBuilder, InserterInner, InserterInnerBuilder,
33 POOL_IDLE_TIMEOUT,
34};
35use super::{
36 Result, SINK_TYPE_APPEND_ONLY, SINK_TYPE_OPTION, SINK_TYPE_UPSERT, SinkError, SinkWriterMetrics,
37};
38use crate::enforce_secret::EnforceSecret;
39use crate::sink::encoder::{DorisJsonConfig, JsonEncoder, RowEncoder};
40use crate::sink::writer::{LogSinkerOf, SinkWriter, SinkWriterExt};
41use crate::sink::{Sink, SinkParam, SinkWriterParam};
42
43pub const DORIS_SINK: &str = "doris";
44
45const fn default_stream_load_http_timeout_ms() -> u64 {
46 30 * 1000
47}
48
49#[derive(Deserialize, Debug, Clone, WithOptions)]
50pub struct DorisCommon {
51 #[serde(rename = "doris.url")]
52 pub url: String,
53 #[serde(rename = "doris.user")]
54 pub user: String,
55 #[serde(rename = "doris.password")]
56 pub password: String,
57 #[serde(rename = "doris.database")]
58 pub database: String,
59 #[serde(rename = "doris.table")]
60 pub table: String,
61 #[serde(rename = "doris.partial_update")]
62 pub partial_update: Option<String>,
63}
64
65impl EnforceSecret for DorisCommon {
66 const ENFORCE_SECRET_PROPERTIES: phf::Set<&'static str> = phf::phf_set! {
67 "doris.password", "doris.user"
68 };
69}
70
71impl DorisCommon {
72 pub(crate) fn build_get_client(&self) -> DorisSchemaClient {
73 DorisSchemaClient::new(
74 self.url.clone(),
75 self.table.clone(),
76 self.database.clone(),
77 self.user.clone(),
78 self.password.clone(),
79 )
80 }
81}
82
83#[serde_as]
84#[derive(Clone, Debug, Deserialize, WithOptions)]
85pub struct DorisConfig {
86 #[serde(flatten)]
87 pub common: DorisCommon,
88
89 pub r#type: String, #[serde(
93 rename = "doris.stream_load.http.timeout.ms",
94 default = "default_stream_load_http_timeout_ms"
95 )]
96 #[serde_as(as = "DisplayFromStr")]
97 #[with_option(allow_alter_on_fly)]
98 pub stream_load_http_timeout_ms: u64,
99
100 #[serde(flatten)]
101 pub unknown_fields: std::collections::HashMap<String, String>,
102}
103
104crate::impl_sink_unknown_fields!(DorisConfig);
105
106impl EnforceSecret for DorisConfig {
107 fn enforce_one(prop: &str) -> crate::error::ConnectorResult<()> {
108 DorisCommon::enforce_one(prop)
109 }
110}
111
112impl DorisConfig {
113 pub fn from_btreemap(properties: BTreeMap<String, String>) -> Result<Self> {
114 let config =
115 serde_json::from_value::<DorisConfig>(serde_json::to_value(properties).unwrap())
116 .map_err(|e| SinkError::Config(anyhow!(e)))?;
117 if config.r#type != SINK_TYPE_APPEND_ONLY && config.r#type != SINK_TYPE_UPSERT {
118 return Err(SinkError::Config(anyhow!(
119 "`{}` must be {}, or {}",
120 SINK_TYPE_OPTION,
121 SINK_TYPE_APPEND_ONLY,
122 SINK_TYPE_UPSERT
123 )));
124 }
125 Ok(config)
126 }
127}
128
129#[derive(Debug)]
130pub struct DorisSink {
131 pub config: DorisConfig,
132 schema: Schema,
133 pk_indices: Vec<usize>,
134 is_append_only: bool,
135}
136
137impl EnforceSecret for DorisSink {
138 fn enforce_secret<'a>(
139 prop_iter: impl Iterator<Item = &'a str>,
140 ) -> crate::error::ConnectorResult<()> {
141 for prop in prop_iter {
142 DorisConfig::enforce_one(prop)?;
143 }
144 Ok(())
145 }
146}
147
148impl DorisSink {
149 pub fn new(
150 config: DorisConfig,
151 schema: Schema,
152 pk_indices: Vec<usize>,
153 is_append_only: bool,
154 ) -> Result<Self> {
155 Ok(Self {
156 config,
157 schema,
158 pk_indices,
159 is_append_only,
160 })
161 }
162}
163
164impl DorisSink {
165 fn check_column_name_and_type(&self, doris_column_fields: Vec<DorisField>) -> Result<()> {
166 let doris_columns_desc: HashMap<String, String> = doris_column_fields
167 .iter()
168 .map(|s| (s.name.clone(), s.r#type.clone()))
169 .collect();
170
171 let rw_fields_name = self.schema.fields();
172 if rw_fields_name.len() > doris_columns_desc.len() {
173 return Err(SinkError::Doris(
174 "The columns of the sink must be equal to or a superset of the target table's columns.".to_owned(),
175 ));
176 }
177
178 for i in rw_fields_name {
179 let value = doris_columns_desc.get(&i.name).ok_or_else(|| {
180 SinkError::Doris(format!(
181 "Column name don't find in doris, risingwave is {:?} ",
182 i.name
183 ))
184 })?;
185 if !Self::check_and_correct_column_type(&i.data_type, value.clone())? {
186 return Err(SinkError::Doris(format!(
187 "Column type don't match, column name is {:?}. doris type is {:?} risingwave type is {:?} ",
188 i.name, value, i.data_type
189 )));
190 }
191 }
192 Ok(())
193 }
194
195 fn check_and_correct_column_type(
196 rw_data_type: &DataType,
197 doris_data_type: String,
198 ) -> Result<bool> {
199 let doris_data_type = doris_data_type.to_ascii_uppercase();
200 let is_variant = doris_data_type.contains("VARIANT");
201 match rw_data_type {
202 risingwave_common::types::DataType::Boolean => Ok(doris_data_type.contains("BOOLEAN")),
203 risingwave_common::types::DataType::Int16 => Ok(doris_data_type.contains("SMALLINT")),
204 risingwave_common::types::DataType::Int32 => Ok(doris_data_type.contains("INT")),
205 risingwave_common::types::DataType::Int64 => Ok(doris_data_type.contains("BIGINT")),
206 risingwave_common::types::DataType::Float32 => Ok(doris_data_type.contains("FLOAT")),
207 risingwave_common::types::DataType::Float64 => Ok(doris_data_type.contains("DOUBLE")),
208 risingwave_common::types::DataType::Decimal => Ok(doris_data_type.contains("DECIMAL")),
209 risingwave_common::types::DataType::Date => Ok(doris_data_type.contains("DATE")),
210 risingwave_common::types::DataType::Varchar => {
211 Ok(
212 doris_data_type.contains("STRING")
213 || doris_data_type.contains("VARCHAR")
214 || is_variant,
215 )
216 }
217 risingwave_common::types::DataType::Time => {
218 Err(SinkError::Doris("TIME is not supported for Doris sink. Please convert to VARCHAR or other supported types.".to_owned()))
219 }
220 risingwave_common::types::DataType::Timestamp => {
221 Ok(doris_data_type.contains("DATETIME"))
222 }
223 risingwave_common::types::DataType::Timestamptz => Err(SinkError::Doris(
224 "TIMESTAMP WITH TIMEZONE is not supported for Doris sink as Doris doesn't store time values with timezone information. Please convert to TIMESTAMP first.".to_owned(),
225 )),
226 risingwave_common::types::DataType::Interval => Err(SinkError::Doris(
227 "INTERVAL is not supported for Doris sink. Please convert to VARCHAR or other supported types.".to_owned(),
228 )),
229 risingwave_common::types::DataType::Struct(_) => Ok(doris_data_type.contains("STRUCT")),
230 risingwave_common::types::DataType::List(_) => Ok(doris_data_type.contains("ARRAY")),
231 risingwave_common::types::DataType::Bytea => {
232 Err(SinkError::Doris("BYTEA is not supported for Doris sink. Please convert to VARCHAR or other supported types.".to_owned()))
233 }
234 risingwave_common::types::DataType::Jsonb => {
235 Ok(doris_data_type.contains("JSON") || is_variant)
236 }
237 risingwave_common::types::DataType::Variant => {
238 Err(SinkError::Doris("VARIANT is not supported for Doris sink.".to_owned()))
239 }
240 risingwave_common::types::DataType::Serial => Ok(doris_data_type.contains("BIGINT")),
241 risingwave_common::types::DataType::Int256 => {
242 Err(SinkError::Doris("INT256 is not supported for Doris sink.".to_owned()))
243 }
244 risingwave_common::types::DataType::Map(_) => {
245 Err(SinkError::Doris("MAP is not supported for Doris sink.".to_owned()))
246 }
247 DataType::Vector(_) => {
248 Err(SinkError::Doris("VECTOR is not supported for Doris sink.".to_owned()))
249 },
250 }
251 }
252}
253
254impl Sink for DorisSink {
255 type LogSinker = LogSinkerOf<DorisSinkWriter>;
256
257 const SINK_NAME: &'static str = DORIS_SINK;
258
259 crate::impl_validate_sink_unknown_fields!();
260
261 async fn new_log_sinker(&self, writer_param: SinkWriterParam) -> Result<Self::LogSinker> {
262 Ok(DorisSinkWriter::new(
263 self.config.clone(),
264 self.schema.clone(),
265 self.pk_indices.clone(),
266 self.is_append_only,
267 )
268 .await?
269 .into_log_sinker(SinkWriterMetrics::new(&writer_param)))
270 }
271
272 async fn validate(&self) -> Result<()> {
273 if !self.is_append_only && self.pk_indices.is_empty() {
274 return Err(SinkError::Config(anyhow!(
275 "Primary key not defined for upsert doris sink (please define in `primary_key` field)"
276 )));
277 }
278 let client = self.config.common.build_get_client();
280 let doris_schema = client.get_schema_from_doris().await?;
281
282 if !self.is_append_only && doris_schema.keys_type.ne("UNIQUE_KEYS") {
283 return Err(SinkError::Config(anyhow!(
284 "If you want to use upsert, please set the keysType of doris to UNIQUE_KEYS"
285 )));
286 }
287 self.check_column_name_and_type(doris_schema.properties)?;
288 Ok(())
289 }
290}
291
292pub struct DorisSinkWriter {
293 pub config: DorisConfig,
294 #[expect(dead_code)]
295 schema: Schema,
296 #[expect(dead_code)]
297 pk_indices: Vec<usize>,
298 inserter_inner_builder: InserterInnerBuilder,
299 is_append_only: bool,
300 client: Option<DorisClient>,
301 row_encoder: JsonEncoder,
302}
303
304impl TryFrom<SinkParam> for DorisSink {
305 type Error = SinkError;
306
307 fn try_from(param: SinkParam) -> std::result::Result<Self, Self::Error> {
308 let schema = param.schema();
309 let pk_indices = param.downstream_pk_or_empty();
310 let config = DorisConfig::from_btreemap(param.properties)?;
311 DorisSink::new(config, schema, pk_indices, param.sink_type.is_append_only())
312 }
313}
314
315impl DorisSinkWriter {
316 pub async fn new(
317 config: DorisConfig,
318 schema: Schema,
319 pk_indices: Vec<usize>,
320 is_append_only: bool,
321 ) -> Result<Self> {
322 let mut decimal_map = HashMap::default();
323 let mut variant_columns = HashSet::default();
324 let doris_schema = config
325 .common
326 .build_get_client()
327 .get_schema_from_doris()
328 .await?;
329 for s in &doris_schema.properties {
330 if let Some(v) = s.get_decimal_pre_scale()? {
331 decimal_map.insert(s.name.clone(), v);
332 }
333 if s.is_variant() {
334 variant_columns.insert(s.name.clone());
335 }
336 }
337
338 let header_builder = HeaderBuilder::new()
339 .add_common_header()
340 .set_user_password(config.common.user.clone(), config.common.password.clone())
341 .add_json_format()
342 .set_partial_columns(config.common.partial_update.clone())
343 .add_read_json_by_line();
344 let header = if !is_append_only {
345 header_builder.add_hidden_column().build()
346 } else {
347 header_builder.build()
348 };
349
350 let doris_insert_builder = InserterInnerBuilder::new(
351 config.common.url.clone(),
352 config.common.database.clone(),
353 config.common.table.clone(),
354 header,
355 config.stream_load_http_timeout_ms,
356 )?;
357 Ok(Self {
358 config,
359 schema: schema.clone(),
360 pk_indices,
361 inserter_inner_builder: doris_insert_builder,
362 is_append_only,
363 client: None,
364 row_encoder: JsonEncoder::new_with_doris(
365 schema,
366 None,
367 DorisJsonConfig {
368 decimal_scale: decimal_map,
369 variant_columns,
370 },
371 ),
372 })
373 }
374
375 async fn append_only(&mut self, chunk: StreamChunk) -> Result<()> {
376 for (op, row) in chunk.rows() {
377 if op != Op::Insert {
378 continue;
379 }
380 let row_json_string = Value::Object(self.row_encoder.encode(row)?).to_string();
381 self.client
382 .as_mut()
383 .ok_or_else(|| SinkError::Doris("Can't find doris sink insert".to_owned()))?
384 .write(row_json_string.into())
385 .await?;
386 }
387 Ok(())
388 }
389
390 async fn upsert(&mut self, chunk: StreamChunk) -> Result<()> {
391 for (op, row) in chunk.rows() {
392 match op {
393 Op::Insert => {
394 let mut row_json_value = self.row_encoder.encode(row)?;
395 row_json_value
396 .insert(DORIS_DELETE_SIGN.to_owned(), Value::String("0".to_owned()));
397 let row_json_string = serde_json::to_string(&row_json_value).map_err(|e| {
398 SinkError::Doris(format!("Json derialize error: {}", e.as_report()))
399 })?;
400 self.client
401 .as_mut()
402 .ok_or_else(|| SinkError::Doris("Can't find doris sink insert".to_owned()))?
403 .write(row_json_string.into())
404 .await?;
405 }
406 Op::Delete => {
407 let mut row_json_value = self.row_encoder.encode(row)?;
408 row_json_value
409 .insert(DORIS_DELETE_SIGN.to_owned(), Value::String("1".to_owned()));
410 let row_json_string = serde_json::to_string(&row_json_value).map_err(|e| {
411 SinkError::Doris(format!("Json derialize error: {}", e.as_report()))
412 })?;
413 self.client
414 .as_mut()
415 .ok_or_else(|| SinkError::Doris("Can't find doris sink insert".to_owned()))?
416 .write(row_json_string.into())
417 .await?;
418 }
419 Op::UpdateDelete => {}
420 Op::UpdateInsert => {
421 let mut row_json_value = self.row_encoder.encode(row)?;
422 row_json_value
423 .insert(DORIS_DELETE_SIGN.to_owned(), Value::String("0".to_owned()));
424 let row_json_string = serde_json::to_string(&row_json_value).map_err(|e| {
425 SinkError::Doris(format!("Json derialize error: {}", e.as_report()))
426 })?;
427 self.client
428 .as_mut()
429 .ok_or_else(|| SinkError::Doris("Can't find doris sink insert".to_owned()))?
430 .write(row_json_string.into())
431 .await?;
432 }
433 }
434 }
435 Ok(())
436 }
437}
438
439#[async_trait]
440impl SinkWriter for DorisSinkWriter {
441 async fn write_batch(&mut self, chunk: StreamChunk) -> Result<()> {
442 if self.client.is_none() {
443 self.client = Some(DorisClient::new(self.inserter_inner_builder.build().await?));
444 }
445 if self.is_append_only {
446 self.append_only(chunk).await
447 } else {
448 self.upsert(chunk).await
449 }
450 }
451
452 async fn begin_epoch(&mut self, _epoch: u64) -> Result<()> {
453 Ok(())
454 }
455
456 async fn abort(&mut self) -> Result<()> {
457 Ok(())
458 }
459
460 async fn barrier(&mut self, _is_checkpoint: bool) -> Result<()> {
461 if self.client.is_some() {
462 let client = self
463 .client
464 .take()
465 .ok_or_else(|| SinkError::Doris("Can't find doris inserter".to_owned()))?;
466 client.finish().await?;
467 }
468 Ok(())
469 }
470}
471
472pub struct DorisSchemaClient {
473 url: String,
474 table: String,
475 db: String,
476 user: String,
477 password: String,
478}
479impl DorisSchemaClient {
480 pub fn new(url: String, table: String, db: String, user: String, password: String) -> Self {
481 Self {
482 url,
483 table,
484 db,
485 user,
486 password,
487 }
488 }
489
490 pub async fn get_schema_from_doris(&self) -> Result<DorisSchema> {
491 let uri = format!("{}/api/{}/{}/_schema", self.url, self.db, self.table);
492
493 let client = reqwest::Client::builder()
494 .pool_idle_timeout(POOL_IDLE_TIMEOUT)
495 .build()
496 .map_err(|err| SinkError::DorisStarrocksConnect(err.into()))?;
497
498 let response = client
499 .get(uri)
500 .header(
501 "Authorization",
502 format!(
503 "Basic {}",
504 general_purpose::STANDARD.encode(format!("{}:{}", self.user, self.password))
505 ),
506 )
507 .send()
508 .await
509 .map_err(|err| SinkError::DorisStarrocksConnect(err.into()))?;
510
511 let json: Value = response
512 .json()
513 .await
514 .map_err(|err| SinkError::DorisStarrocksConnect(err.into()))?;
515 let json_data = if json.get("code").is_some() && json.get("msg").is_some() {
516 json.get("data")
517 .ok_or_else(|| {
518 SinkError::DorisStarrocksConnect(anyhow::anyhow!("Can't find data"))
519 })?
520 .clone()
521 } else {
522 json
523 };
524 let schema: DorisSchema = serde_json::from_value(json_data)
525 .context("Can't get schema from json")
526 .map_err(SinkError::DorisStarrocksConnect)?;
527 Ok(schema)
528 }
529}
530#[derive(Debug, Serialize, Deserialize)]
531pub struct DorisSchema {
532 status: i32,
533 #[serde(rename = "keysType")]
534 pub keys_type: String,
535 pub properties: Vec<DorisField>,
536}
537#[derive(Debug, Serialize, Deserialize)]
538pub struct DorisField {
539 pub name: String,
540 pub r#type: String,
541 comment: String,
542 pub precision: Option<String>,
543 pub scale: Option<String>,
544 aggregation_type: String,
545}
546impl DorisField {
547 pub fn get_decimal_pre_scale(&self) -> Result<Option<u8>> {
548 if self.r#type.contains("DECIMAL") {
549 let scale = self
550 .scale
551 .as_ref()
552 .ok_or_else(|| {
553 SinkError::Doris(format!(
554 "In doris, the type of {} is DECIMAL, but `scale` is not found",
555 self.name
556 ))
557 })?
558 .parse::<u8>()
559 .map_err(|err| {
560 SinkError::Doris(format!(
561 "Unable to convert decimal's scale to u8. error: {:?}",
562 err.kind()
563 ))
564 })?;
565 Ok(Some(scale))
566 } else {
567 Ok(None)
568 }
569 }
570
571 pub fn is_variant(&self) -> bool {
572 self.r#type.to_ascii_uppercase().contains("VARIANT")
573 }
574}
575
576#[cfg(test)]
577mod tests {
578 use risingwave_common::types::DataType;
579
580 use super::DorisSink;
581
582 #[test]
583 fn test_jsonb_can_write_to_variant() {
584 assert!(
585 DorisSink::check_and_correct_column_type(&DataType::Jsonb, "VARIANT".into()).unwrap()
586 );
587 }
588
589 #[test]
590 fn test_varchar_can_write_to_variant() {
591 assert!(
592 DorisSink::check_and_correct_column_type(&DataType::Varchar, "VARIANT".into()).unwrap()
593 );
594 }
595}
596
597#[derive(Debug, Serialize, Deserialize)]
598pub struct DorisInsertResultResponse {
599 #[serde(rename = "TxnId")]
600 txn_id: i64,
601 #[serde(rename = "Label")]
602 label: String,
603 #[serde(rename = "Status")]
604 status: String,
605 #[serde(rename = "TwoPhaseCommit")]
606 two_phase_commit: String,
607 #[serde(rename = "Message")]
608 message: String,
609 #[serde(rename = "NumberTotalRows")]
610 number_total_rows: i64,
611 #[serde(rename = "NumberLoadedRows")]
612 number_loaded_rows: i64,
613 #[serde(rename = "NumberFilteredRows")]
614 number_filtered_rows: i32,
615 #[serde(rename = "NumberUnselectedRows")]
616 number_unselected_rows: i32,
617 #[serde(rename = "LoadBytes")]
618 load_bytes: i64,
619 #[serde(rename = "LoadTimeMs")]
620 load_time_ms: i32,
621 #[serde(rename = "BeginTxnTimeMs")]
622 begin_txn_time_ms: i32,
623 #[serde(rename = "StreamLoadPutTimeMs")]
624 stream_load_put_time_ms: i32,
625 #[serde(rename = "ReadDataTimeMs")]
626 read_data_time_ms: i32,
627 #[serde(rename = "WriteDataTimeMs")]
628 write_data_time_ms: i32,
629 #[serde(rename = "CommitAndPublishTimeMs")]
630 commit_and_publish_time_ms: i32,
631 #[serde(rename = "ErrorURL")]
632 err_url: Option<String>,
633}
634
635pub struct DorisClient {
636 insert: InserterInner,
637 is_first_record: bool,
638}
639impl DorisClient {
640 pub fn new(insert: InserterInner) -> Self {
641 Self {
642 insert,
643 is_first_record: true,
644 }
645 }
646
647 pub async fn write(&mut self, data: Bytes) -> Result<()> {
648 let mut data_build = BytesMut::new();
649 if self.is_first_record {
650 self.is_first_record = false;
651 } else {
652 data_build.put_slice("\n".as_bytes());
653 }
654 data_build.put_slice(&data);
655 self.insert.write(data_build.into()).await?;
656 Ok(())
657 }
658
659 pub async fn finish(self) -> Result<DorisInsertResultResponse> {
660 let raw = self.insert.finish().await?;
661 let res: DorisInsertResultResponse = serde_json::from_slice(&raw)
662 .map_err(|err| SinkError::DorisStarrocksConnect(err.into()))?;
663
664 if !DORIS_SUCCESS_STATUS.contains(&res.status.as_str()) {
665 return Err(SinkError::DorisStarrocksConnect(anyhow::anyhow!(
666 "Insert error: {:?}, error url: {:?}",
667 res.message,
668 res.err_url
669 )));
670 };
671 Ok(res)
672 }
673}