1use std::sync::Arc;
49use std::sync::atomic::{AtomicU32, Ordering};
50use std::time::{Duration, SystemTime, UNIX_EPOCH};
51
52use axum::Router;
53use axum::extract::ws::{Message, WebSocket, WebSocketUpgrade};
54use axum::extract::{Extension, Path};
55use axum::http::{HeaderMap, StatusCode};
56use axum::response::IntoResponse;
57use axum::routing::get;
58use futures::SinkExt;
59use futures::stream::StreamExt;
60use jsonbb::Value;
61use risingwave_common::array::{Op, StreamChunk};
62use risingwave_common::license::Feature;
63use risingwave_common::row::OwnedRow;
64use risingwave_common::types::{DataType, JsonbVal};
65use risingwave_common::util::chunk_coalesce::DataChunkBuilder;
66use risingwave_pb::task_service::{
67 IngestDmlInitRequest, IngestDmlPayloadRequest, IngestDmlRequest, ingest_dml_request,
68 ingest_dml_response,
69};
70use serde::{Deserialize, Serialize};
71use serde_json::value::RawValue;
72use thiserror_ext::AsReport;
73use tokio::time::timeout;
74use tokio_stream::wrappers::ReceiverStream;
75use tower::ServiceBuilder;
76use tower_http::add_extension::AddExtensionLayer;
77
78use crate::session::SESSION_MANAGER;
79use crate::webhook::payload::{build_json_access_builder, owned_row_from_payload_row};
80use crate::webhook::utils::{authenticate_webhook_payload, err, header_map_to_json};
81use crate::webhook::{PayloadSchema, acquire_table_info};
82
83const INIT_MESSAGE_TYPE: &str = "init";
84const INGEST_DML_REQUEST_BUFFER_SIZE: usize = 64;
85const WEBSOCKET_INIT_TIMEOUT: Duration = Duration::from_secs(10);
86
87pub struct IngestService {
89 counter: AtomicU32,
90}
91
92impl IngestService {
93 pub fn new() -> Self {
94 Self {
95 counter: AtomicU32::new(0),
96 }
97 }
98}
99
100pub type ServiceRef = Arc<IngestService>;
101
102#[derive(Debug, Deserialize)]
103struct InitRequest {
104 #[serde(rename = "type")]
105 msg_type: String,
106 timestamp: i64,
107}
108
109#[derive(Debug, Deserialize)]
110struct RawDmlBatchRequest {
111 pub dml_batch_id: u64,
112 pub items: Vec<RawDmlRequest>,
113}
114
115#[derive(Debug, Deserialize)]
116pub struct RawDmlRequest {
117 pub op: Option<String>,
118 pub data: Box<RawValue>,
120}
121
122#[derive(Debug, Clone, Copy)]
123enum DmlOp {
124 Upsert,
125 Delete,
126}
127
128#[derive(Debug)]
129pub struct DmlRequest {
130 op: DmlOp,
131 data: Box<RawValue>,
132}
133
134impl TryFrom<RawDmlRequest> for DmlRequest {
135 type Error = String;
136
137 fn try_from(raw: RawDmlRequest) -> Result<Self, Self::Error> {
138 let op = match raw.op.as_deref() {
139 None => {
140 return Err(
141 "missing `op`; expected `upsert`, `insert`, `update`, or `delete`".to_owned(),
142 );
143 }
144 Some("upsert" | "insert" | "update") => DmlOp::Upsert,
145 Some("delete") => DmlOp::Delete,
146 Some(other) => {
147 return Err(format!(
148 "unknown `op` value `{other}`; expected `upsert`, `insert`, `update`, or `delete`"
149 ));
150 }
151 };
152
153 Ok(Self { op, data: raw.data })
154 }
155}
156
157#[derive(Debug, Serialize)]
158#[serde(untagged)]
159pub enum ServerMessage {
160 Ack { ack: u64 },
161 Fatal { fatal: String },
162}
163
164#[derive(Debug)]
165struct PreparedDmlBatch {
166 dml_batch_id: u64,
167 payload: IngestDmlPayloadRequest,
168}
169
170type WsTx = futures::stream::SplitSink<WebSocket, Message>;
171type WsRx = futures::stream::SplitStream<WebSocket>;
172
173fn check_websocket_ingest_license() -> crate::webhook::utils::Result<()> {
175 Feature::WebSocketIngest
176 .check_available()
177 .map_err(|e| err(e, StatusCode::FORBIDDEN))
178}
179
180pub async fn ws_handler(
181 ws: WebSocketUpgrade,
182 Extension(svc): Extension<ServiceRef>,
183 headers: HeaderMap,
184 Path((database, schema, table)): Path<(String, String, String)>,
185) -> crate::webhook::utils::Result<impl IntoResponse> {
186 check_websocket_ingest_license()?;
187
188 let request_id = svc.counter.fetch_add(1, Ordering::Relaxed);
189 let headers_jsonb = header_map_to_json(&headers);
190 Ok(ws.on_upgrade(move |socket| {
191 handle_connection(
192 socket,
193 database,
194 schema,
195 table,
196 request_id,
197 headers,
198 headers_jsonb,
199 )
200 }))
201}
202
203async fn handle_connection(
204 socket: WebSocket,
205 database: String,
206 schema: String,
207 table: String,
208 request_id: u32,
209 headers: HeaderMap,
210 headers_jsonb: JsonbVal,
211) {
212 let (mut ws_tx, ws_rx) = socket.split();
213
214 if let Err(e) = try_handle_connection(
215 &mut ws_tx,
216 ws_rx,
217 database,
218 schema,
219 table,
220 request_id,
221 headers,
222 headers_jsonb,
223 )
224 .await
225 {
226 let _ = send_fatal(&mut ws_tx, e).await;
227 }
228}
229
230async fn try_handle_connection(
231 ws_tx: &mut WsTx,
232 mut ws_rx: WsRx,
233 database: String,
234 schema: String,
235 table: String,
236 request_id: u32,
237 headers: HeaderMap,
238 headers_jsonb: JsonbVal,
239) -> Result<(), String> {
240 let table_info = acquire_table_info(request_id, &database, &schema, &table)
241 .await
242 .map_err(|e| format!("table lookup failed: {}", e.as_report()))?;
243
244 let session_mgr = SESSION_MANAGER.get().expect("session manager initialized");
245 let max_clock_skew_ms = session_mgr
246 .env()
247 .frontend_config()
248 .webhook_auth_max_clock_skew_ms;
249 let webhook_source_info = table_info.webhook_source_info;
250 let table_id = table_info.table_id;
251 let table_version_id = table_info.table_version_id;
252 let row_id_index = table_info.row_id_index;
253 let compute_client = table_info.compute_client;
254 let payload_schema = table_info.payload_schema;
255
256 let init_text = match timeout(WEBSOCKET_INIT_TIMEOUT, ws_rx.next()).await {
257 Ok(Some(Ok(Message::Text(text)))) => text,
258 Ok(Some(Ok(Message::Close(_)))) | Ok(None) => return Ok(()),
259 Ok(Some(Ok(_))) => {
260 return Err("the first WebSocket frame must be a text `init` message".to_owned());
261 }
262 Ok(Some(Err(e))) => {
263 return Err(format!(
264 "failed to read WebSocket init message: {}",
265 e.as_report()
266 ));
267 }
268 Err(_) => {
269 return Err(format!(
270 "timed out waiting for WebSocket init message after {}s",
271 WEBSOCKET_INIT_TIMEOUT.as_secs()
272 ));
273 }
274 };
275
276 if let Err(e) =
277 authenticate_webhook_payload(headers_jsonb, init_text.as_bytes(), &webhook_source_info)
278 .await
279 {
280 return Err(format!("{}", e.as_report()));
281 }
282
283 parse_and_validate_init_request(&init_text, max_clock_skew_ms)?;
284
285 let (ingest_req_tx, ingest_req_rx) = tokio::sync::mpsc::channel(INGEST_DML_REQUEST_BUFFER_SIZE);
286 if ingest_req_tx
287 .send(IngestDmlRequest {
288 request: Some(ingest_dml_request::Request::Init(IngestDmlInitRequest {
289 table_id,
290 table_version_id,
291 request_id,
292 row_id_index,
293 })),
294 })
295 .await
296 .is_err()
297 {
298 return Err("failed to enqueue init request for ingest stream".to_owned());
299 }
300
301 let mut ingest_resp_stream = compute_client
302 .ingest_dml(ReceiverStream::new(ingest_req_rx))
303 .await
304 .map_err(|e| format!("failed to open ingest stream: {}", e.as_report()))?;
305
306 match ingest_resp_stream.message().await {
307 Ok(Some(resp)) => match resp.response {
308 Some(ingest_dml_response::Response::Init(_)) => {}
309 _ => {
310 return Err("unexpected init response from ingest stream".to_owned());
311 }
312 },
313 Ok(None) => return Err("ingest stream closed during init".to_owned()),
314 Err(e) => {
315 return Err(format!("ingest stream init error: {}", e.as_report()));
316 }
317 }
318
319 let mut last_seen_dml_batch_id = 0_u64;
320 let mut last_forwarded_dml_batch_id = 0_u64;
321 let mut last_acked_dml_batch_id = 0_u64;
322
323 loop {
324 tokio::select! {
325 biased;
326 ingest_resp = ingest_resp_stream.message() => {
327 match ingest_resp {
328 Ok(Some(resp)) => match resp.response {
329 Some(ingest_dml_response::Response::Ack(ack)) => {
330 if ack.dml_batch_id <= last_acked_dml_batch_id
331 || ack.dml_batch_id > last_forwarded_dml_batch_id
332 {
333 return Err(format!(
334 "received an unexpected ack for dml_batch_id {}",
335 ack.dml_batch_id
336 ));
337 }
338 last_acked_dml_batch_id = ack.dml_batch_id;
339 send_server_message(
340 ws_tx,
341 ServerMessage::Ack {
342 ack: ack.dml_batch_id,
343 },
344 )
345 .await
346 .map_err(|_| "WebSocket connection closed while sending an ack".to_owned())?;
347 }
348 Some(ingest_dml_response::Response::Init(_)) => {
349 return Err("unexpected extra init response from ingest stream".to_owned());
350 }
351 None => return Err("empty response from ingest stream".to_owned()),
352 },
353 Ok(None) => return Err("ingest stream closed".to_owned()),
354 Err(e) => return Err(format!("ingest stream error: {}", e.as_report())),
355 }
356 }
357
358 ws_msg = ws_rx.next() => {
359 let text = match ws_msg {
360 Some(Ok(Message::Text(text))) => text,
361 Some(Ok(Message::Close(_))) | None => break,
362 Some(Ok(_)) => continue,
363 Some(Err(e)) => {
364 return Err(format!(
365 "failed to read WebSocket message: {}",
366 e.as_report()
367 ));
368 }
369 };
370
371 let raw_dml_batch = match serde_json::from_str::<RawDmlBatchRequest>(&text) {
372 Ok(batch) => batch,
373 Err(e) => return Err(format!("malformed payload: {}", e.as_report())),
374 };
375
376 last_seen_dml_batch_id = validate_monotonic_dml_batch_id(
377 raw_dml_batch.dml_batch_id,
378 last_seen_dml_batch_id,
379 )?;
380
381 if raw_dml_batch.items.is_empty() {
382 send_server_message(
383 ws_tx,
384 ServerMessage::Ack {
385 ack: raw_dml_batch.dml_batch_id,
386 },
387 )
388 .await
389 .map_err(|_| "WebSocket connection closed while sending an ack".to_owned())?;
390 continue;
391 }
392
393 let prepared_batch = prepare_dml_batch_payload(
394 &headers,
395 raw_dml_batch,
396 &payload_schema,
397 )
398 .map_err(|e| format!("failed to prepare DML batch: {e}"))?;
399
400 let PreparedDmlBatch {
401 dml_batch_id,
402 payload,
403 } = prepared_batch;
404
405 ingest_req_tx
406 .send(IngestDmlRequest {
407 request: Some(ingest_dml_request::Request::Payload(payload)),
408 })
409 .await
410 .map_err(|_| "ingest stream request channel closed".to_owned())?;
411 last_forwarded_dml_batch_id = dml_batch_id;
412 }
413 }
414 }
415
416 Ok(())
417}
418
419fn parse_and_validate_init_request(
420 text: &str,
421 max_clock_skew_ms: u64,
422) -> Result<InitRequest, String> {
423 let init_req: InitRequest = serde_json::from_str(text)
424 .map_err(|e| format!("malformed init message: {}", e.as_report()))?;
425
426 if init_req.msg_type != INIT_MESSAGE_TYPE {
427 return Err(format!(
428 "invalid init message type '{}', expected '{}'",
429 init_req.msg_type, INIT_MESSAGE_TYPE
430 ));
431 }
432
433 validate_timestamp_skew(init_req.timestamp, max_clock_skew_ms)?;
434 Ok(init_req)
435}
436
437fn validate_timestamp_skew(timestamp_ms: i64, max_clock_skew_ms: u64) -> Result<(), String> {
438 if timestamp_ms < 0 {
439 return Err("timestamp must be a non-negative Unix timestamp in milliseconds".to_owned());
440 }
441
442 let now_ms = SystemTime::now()
443 .duration_since(UNIX_EPOCH)
444 .unwrap_or(Duration::ZERO)
445 .as_millis() as i128;
446 let diff_ms = (now_ms - i128::from(timestamp_ms)).abs();
447
448 if diff_ms > i128::from(max_clock_skew_ms) {
449 return Err(format!(
450 "timestamp skew {}ms exceeds the allowed {}ms window",
451 diff_ms, max_clock_skew_ms
452 ));
453 }
454
455 Ok(())
456}
457
458fn validate_monotonic_dml_batch_id(
459 dml_batch_id: u64,
460 last_seen_dml_batch_id: u64,
461) -> Result<u64, String> {
462 if dml_batch_id <= last_seen_dml_batch_id {
463 return Err(format!(
464 "dml_batch_id must increase monotonically: received {} after {}",
465 dml_batch_id, last_seen_dml_batch_id
466 ));
467 }
468 Ok(dml_batch_id)
469}
470
471fn prepare_dml_batch_payload(
472 headers: &HeaderMap,
473 raw_dml_batch: RawDmlBatchRequest,
474 payload_schema: &PayloadSchema,
475) -> Result<PreparedDmlBatch, String> {
476 let dml_batch_id = raw_dml_batch.dml_batch_id;
477 let raw_dml_reqs = raw_dml_batch.items;
478
479 match payload_schema {
480 PayloadSchema::SingleJsonb => {
481 let mut chunk_builder = DataChunkBuilder::new(
482 vec![DataType::Jsonb],
483 raw_dml_reqs.len().saturating_add(1).max(1),
484 );
485 let mut ops = Vec::with_capacity(raw_dml_reqs.len());
486
487 for (index, raw_dml_req) in raw_dml_reqs.into_iter().enumerate() {
488 let item_index = index + 1;
489 let dml_req = DmlRequest::try_from(raw_dml_req)
490 .map_err(|e| format!("dml_batch_id {dml_batch_id} item {item_index}: {e}"))?;
491
492 let row = Value::from_text(dml_req.data.get().as_bytes())
493 .map(|json_value| OwnedRow::new(vec![Some(JsonbVal::from(json_value).into())]))
494 .map_err(|e| {
495 format!(
496 "dml_batch_id {dml_batch_id} item {item_index}: failed to parse request body: {}",
497 e.as_report()
498 )
499 })?;
500
501 let output = chunk_builder.append_one_row(row);
502 debug_assert!(output.is_none());
503 ops.push(match dml_req.op {
504 DmlOp::Upsert => Op::Insert,
505 DmlOp::Delete => Op::Delete,
506 });
507 }
508
509 let data_chunk = chunk_builder
510 .consume_all()
511 .expect("buffered rows should produce a chunk");
512 let chunk = StreamChunk::from_parts(ops, data_chunk);
513 let payload = IngestDmlPayloadRequest {
514 dml_batch_id,
515 chunk: Some(chunk.to_protobuf()),
516 };
517
518 Ok(PreparedDmlBatch {
519 dml_batch_id,
520 payload,
521 })
522 }
523 PayloadSchema::FullSchema { columns } => {
524 let mut access_builder =
525 build_json_access_builder(headers).map_err(|e| format!("{}", e.as_report()))?;
526 let mut chunk_builder = DataChunkBuilder::new(
527 columns
528 .iter()
529 .map(|column| column.data_type.clone())
530 .collect(),
531 raw_dml_reqs.len().saturating_add(1).max(1),
532 );
533 let mut ops = Vec::with_capacity(raw_dml_reqs.len());
534
535 for (index, raw_dml_req) in raw_dml_reqs.into_iter().enumerate() {
536 let item_index = index + 1;
537 let dml_req = DmlRequest::try_from(raw_dml_req)
538 .map_err(|e| format!("dml_batch_id {dml_batch_id} item {item_index}: {e}"))?;
539
540 let row = owned_row_from_payload_row(
541 &mut access_builder,
542 columns,
543 dml_req.data.get().as_bytes(),
544 )
545 .map_err(|e| {
546 format!(
547 "dml_batch_id {dml_batch_id} item {item_index}: {}",
548 e.as_report()
549 )
550 })?;
551
552 let output = chunk_builder.append_one_row(row);
553 debug_assert!(output.is_none());
554 ops.push(match dml_req.op {
555 DmlOp::Upsert => Op::Insert,
556 DmlOp::Delete => Op::Delete,
557 });
558 }
559
560 let data_chunk = chunk_builder
561 .consume_all()
562 .expect("buffered rows should produce a chunk");
563 let chunk = StreamChunk::from_parts(ops, data_chunk);
564 let payload = IngestDmlPayloadRequest {
565 dml_batch_id,
566 chunk: Some(chunk.to_protobuf()),
567 };
568
569 Ok(PreparedDmlBatch {
570 dml_batch_id,
571 payload,
572 })
573 }
574 }
575}
576
577async fn send_server_message(ws_tx: &mut WsTx, msg: ServerMessage) -> Result<(), String> {
578 let text = serde_json::to_string(&msg).map_err(|e| format!("{}", e.as_report()))?;
579 ws_tx
580 .send(Message::Text(text.into()))
581 .await
582 .map_err(|e| format!("{}", e.as_report()))
583}
584
585async fn send_fatal(ws_tx: &mut WsTx, fatal: String) -> Result<(), String> {
586 send_server_message(ws_tx, ServerMessage::Fatal { fatal }).await
587}
588
589pub fn build_router(svc: ServiceRef) -> Router {
590 Router::new()
591 .route("/{database}/{schema}/{table}", get(ws_handler))
592 .layer(
593 ServiceBuilder::new()
594 .layer(AddExtensionLayer::new(svc))
595 .into_inner(),
596 )
597}
598
599#[cfg(test)]
600mod tests {
601 use axum::http::{HeaderMap, HeaderValue};
602 use risingwave_common::license::{LicenseKey, LicenseManager};
603 use risingwave_common::row::Row;
604 use risingwave_common::types::{DataType, ScalarImpl, ToOwnedDatum};
605
606 use super::*;
607 use crate::webhook::WebhookTableColumnDesc;
608
609 fn raw_json(text: &str) -> Box<RawValue> {
610 serde_json::from_str(text).unwrap()
611 }
612
613 fn raw_req(op: Option<&str>, data: &str) -> RawDmlRequest {
614 RawDmlRequest {
615 op: op.map(str::to_owned),
616 data: raw_json(data),
617 }
618 }
619
620 fn raw_batch(dml_batch_id: u64, items: Vec<RawDmlRequest>) -> RawDmlBatchRequest {
621 RawDmlBatchRequest {
622 dml_batch_id,
623 items,
624 }
625 }
626
627 struct RestoreDefaultLicense;
628
629 impl Drop for RestoreDefaultLicense {
630 fn drop(&mut self) {
631 LicenseManager::get().refresh(LicenseKey::default().as_ref());
632 }
633 }
634
635 #[test]
636 fn test_websocket_ingest_requires_license() {
637 let _restore = RestoreDefaultLicense;
638 LicenseManager::get().refresh(LicenseKey::default().as_ref());
639 check_websocket_ingest_license().unwrap();
640
641 LicenseManager::get().refresh(LicenseKey::empty().as_ref());
642
643 let err = check_websocket_ingest_license().unwrap_err();
644
645 assert_eq!(err.code(), StatusCode::FORBIDDEN);
646 assert!(err.to_string().contains("WebSocketIngest"));
647 }
648
649 fn test_columns(columns: &[(&str, DataType, bool)]) -> Vec<WebhookTableColumnDesc> {
650 columns
651 .iter()
652 .map(|(name, data_type, is_pk)| WebhookTableColumnDesc {
653 name: (*name).to_owned(),
654 data_type: data_type.clone(),
655 is_pk: *is_pk,
656 })
657 .collect()
658 }
659
660 #[test]
661 fn test_validate_monotonic_dml_batch_id() {
662 assert_eq!(validate_monotonic_dml_batch_id(5, 1).unwrap(), 5);
663 }
664
665 #[test]
666 fn test_validate_monotonic_dml_batch_id_rejects_non_monotonic_sequence() {
667 let err = validate_monotonic_dml_batch_id(3, 3).unwrap_err();
668 assert!(err.contains("dml_batch_id must increase monotonically"));
669 }
670
671 #[test]
672 fn test_parse_and_validate_init_request() {
673 let now_ms = SystemTime::now()
674 .duration_since(UNIX_EPOCH)
675 .unwrap_or(Duration::ZERO)
676 .as_millis() as i64;
677 let init = format!(r#"{{"type":"init","timestamp":{now_ms}}}"#);
678 assert_eq!(
679 parse_and_validate_init_request(&init, 300_000)
680 .unwrap()
681 .timestamp,
682 now_ms
683 );
684 }
685
686 #[test]
687 fn test_parse_and_validate_init_request_rejects_stale_timestamp() {
688 let init = r#"{"type":"init","timestamp":0}"#;
689 let err = parse_and_validate_init_request(init, 1).unwrap_err();
690 assert!(err.contains("timestamp skew"));
691 }
692
693 #[test]
694 fn test_dml_request_requires_op() {
695 let err = DmlRequest::try_from(raw_req(None, r#"{"id":1}"#)).unwrap_err();
696 assert_eq!(
697 err,
698 "missing `op`; expected `upsert`, `insert`, `update`, or `delete`"
699 );
700 }
701
702 #[test]
703 fn test_empty_batch_is_valid_json_protocol_input() {
704 let raw_dml_batch = raw_batch(42, vec![]);
705 assert_eq!(
706 validate_monotonic_dml_batch_id(raw_dml_batch.dml_batch_id, 41).unwrap(),
707 42
708 );
709 assert!(raw_dml_batch.items.is_empty());
710 }
711
712 #[test]
713 fn test_prepare_dml_batch_payload_builds_single_chunk_for_batch() {
714 let raw_dml_batch = raw_batch(
715 10,
716 vec![
717 raw_req(
718 Some("upsert"),
719 r#"{"id":1,"price":"19.99","created_at":"2026-04-15 10:00:00","name":"alice"}"#,
720 ),
721 raw_req(
722 Some("delete"),
723 r#"{"id":2,"price":"29.99","created_at":"2026-04-16 10:00:00","name":"bob"}"#,
724 ),
725 ],
726 );
727 let payload_schema = PayloadSchema::FullSchema {
728 columns: test_columns(&[
729 ("id", DataType::Int32, true),
730 ("price", DataType::Decimal, false),
731 ("created_at", DataType::Timestamp, false),
732 ("name", DataType::Varchar, false),
733 ]),
734 };
735
736 let prepared_batch =
737 prepare_dml_batch_payload(&HeaderMap::new(), raw_dml_batch, &payload_schema).unwrap();
738
739 assert_eq!(prepared_batch.dml_batch_id, 10);
740
741 let payload = prepared_batch.payload;
742 assert_eq!(payload.dml_batch_id, 10);
743 let chunk = StreamChunk::from_protobuf(payload.chunk.as_ref().unwrap()).unwrap();
744
745 assert_eq!(chunk.ops(), &[Op::Insert, Op::Delete]);
746 let mut rows = chunk.rows();
747 let row = rows.next().unwrap().1;
748 assert_eq!(row.datum_at(0).to_owned_datum(), Some(ScalarImpl::Int32(1)));
749 assert!(matches!(
750 row.datum_at(1).to_owned_datum(),
751 Some(ScalarImpl::Decimal(_))
752 ));
753 assert!(matches!(
754 row.datum_at(2).to_owned_datum(),
755 Some(ScalarImpl::Timestamp(_))
756 ));
757 assert_eq!(
758 row.datum_at(3).to_owned_datum(),
759 Some(ScalarImpl::Utf8("alice".into()))
760 );
761
762 let row = rows.next().unwrap().1;
763 assert_eq!(row.datum_at(0).to_owned_datum(), Some(ScalarImpl::Int32(2)));
764 }
765
766 #[test]
767 fn test_prepare_dml_batch_payload_returns_first_item_error() {
768 let raw_dml_batch = raw_batch(
769 11,
770 vec![
771 raw_req(Some("delete"), r#"{"id":1,"name":"alice"}"#),
772 raw_req(Some("upsert"), r#"{"name":"bob"}"#),
773 ],
774 );
775 let payload_schema = PayloadSchema::FullSchema {
776 columns: test_columns(&[
777 ("id", DataType::Int32, true),
778 ("name", DataType::Varchar, false),
779 ]),
780 };
781
782 let err = prepare_dml_batch_payload(&HeaderMap::new(), raw_dml_batch, &payload_schema)
783 .unwrap_err();
784
785 assert!(err.contains("dml_batch_id 11 item 2"));
786 assert!(err.contains("failed to decode webhook JSON payload"));
787 }
788
789 #[test]
790 fn test_prepare_dml_batch_payload_allows_delete_with_only_pk() {
791 let raw_dml_batch = raw_batch(
792 12,
793 vec![
794 raw_req(Some("upsert"), r#"{"id":7,"name":"alice"}"#),
795 raw_req(Some("delete"), r#"{"id":7}"#),
796 ],
797 );
798 let payload_schema = PayloadSchema::FullSchema {
799 columns: test_columns(&[
800 ("id", DataType::Int32, true),
801 ("name", DataType::Varchar, false),
802 ]),
803 };
804
805 let prepared_batch =
806 prepare_dml_batch_payload(&HeaderMap::new(), raw_dml_batch, &payload_schema).unwrap();
807 let chunk =
808 StreamChunk::from_protobuf(prepared_batch.payload.chunk.as_ref().unwrap()).unwrap();
809
810 assert_eq!(prepared_batch.dml_batch_id, 12);
811 assert_eq!(chunk.ops(), &[Op::Insert, Op::Delete]);
812 let mut rows = chunk.rows();
813 assert_eq!(
814 rows.next().unwrap().1.datum_at(1).to_owned_datum(),
815 Some(ScalarImpl::Utf8("alice".into()))
816 );
817 assert_eq!(rows.next().unwrap().1.datum_at(1).to_owned_datum(), None);
818 }
819
820 #[test]
821 fn test_prepare_dml_batch_payload_rejects_incomplete_composite_pk_insert() {
822 let raw_dml_batch = raw_batch(
823 51,
824 vec![raw_req(Some("upsert"), r#"{"id":1,"name":"alice"}"#)],
825 );
826 let payload_schema = PayloadSchema::FullSchema {
827 columns: test_columns(&[
828 ("tenant_id", DataType::Int32, true),
829 ("id", DataType::Int32, true),
830 ("name", DataType::Varchar, false),
831 ]),
832 };
833
834 let err = prepare_dml_batch_payload(&HeaderMap::new(), raw_dml_batch, &payload_schema)
835 .unwrap_err();
836
837 assert!(err.contains("dml_batch_id 51 item 1"));
838 assert!(err.contains("failed to decode webhook JSON payload"));
839 }
840
841 #[test]
842 fn test_prepare_dml_batch_payload_rejects_incomplete_composite_pk_delete() {
843 let raw_dml_batch = raw_batch(52, vec![raw_req(Some("delete"), r#"{"tenant_id":1}"#)]);
844 let payload_schema = PayloadSchema::FullSchema {
845 columns: test_columns(&[
846 ("tenant_id", DataType::Int32, true),
847 ("id", DataType::Int32, true),
848 ("name", DataType::Varchar, false),
849 ]),
850 };
851
852 let err = prepare_dml_batch_payload(&HeaderMap::new(), raw_dml_batch, &payload_schema)
853 .unwrap_err();
854
855 assert!(err.contains("dml_batch_id 52 item 1"));
856 assert!(err.contains("failed to decode webhook JSON payload"));
857 }
858
859 #[test]
860 fn test_prepare_dml_batch_payload_applies_supported_decoder_headers() {
861 let mut headers = HeaderMap::new();
862 headers.insert(
863 "x-rw-webhook-json-timestamp-handling-mode",
864 HeaderValue::from_static("milli"),
865 );
866 headers.insert(
867 "x-rw-webhook-json-time-handling-mode",
868 HeaderValue::from_static("milli"),
869 );
870 headers.insert(
871 "x-rw-webhook-json-bigint-unsigned-handling-mode",
872 HeaderValue::from_static("precise"),
873 );
874
875 let payload_schema = PayloadSchema::FullSchema {
876 columns: test_columns(&[
877 ("id", DataType::Int32, true),
878 ("event_time", DataType::Timestamp, false),
879 ]),
880 };
881 let prepared_timestamp_batch = prepare_dml_batch_payload(
882 &headers,
883 raw_batch(
884 13,
885 vec![raw_req(
886 Some("upsert"),
887 r#"{"id":1,"event_time":1712800800123}"#,
888 )],
889 ),
890 &payload_schema,
891 )
892 .unwrap();
893 let timestamp_chunk =
894 StreamChunk::from_protobuf(prepared_timestamp_batch.payload.chunk.as_ref().unwrap())
895 .unwrap();
896 assert!(matches!(
897 timestamp_chunk
898 .rows()
899 .next()
900 .unwrap()
901 .1
902 .datum_at(1)
903 .to_owned_datum(),
904 Some(ScalarImpl::Timestamp(_))
905 ));
906
907 let payload_schema = PayloadSchema::FullSchema {
908 columns: test_columns(&[
909 ("id", DataType::Int32, true),
910 ("event_time", DataType::Time, false),
911 ]),
912 };
913 let prepared_time_batch = prepare_dml_batch_payload(
914 &headers,
915 raw_batch(
916 14,
917 vec![raw_req(Some("upsert"), r#"{"id":2,"event_time":3723123}"#)],
918 ),
919 &payload_schema,
920 )
921 .unwrap();
922 let time_chunk =
923 StreamChunk::from_protobuf(prepared_time_batch.payload.chunk.as_ref().unwrap())
924 .unwrap();
925 assert!(matches!(
926 time_chunk
927 .rows()
928 .next()
929 .unwrap()
930 .1
931 .datum_at(1)
932 .to_owned_datum(),
933 Some(ScalarImpl::Time(_))
934 ));
935
936 let payload_schema = PayloadSchema::FullSchema {
937 columns: test_columns(&[
938 ("id", DataType::Int32, true),
939 ("amount", DataType::Decimal, false),
940 ]),
941 };
942 let prepared_decimal_batch = prepare_dml_batch_payload(
943 &headers,
944 raw_batch(
945 15,
946 vec![raw_req(Some("upsert"), r#"{"id":3,"amount":"AeJA"}"#)],
947 ),
948 &payload_schema,
949 )
950 .unwrap();
951 let decimal_chunk =
952 StreamChunk::from_protobuf(prepared_decimal_batch.payload.chunk.as_ref().unwrap())
953 .unwrap();
954 assert!(matches!(
955 decimal_chunk
956 .rows()
957 .next()
958 .unwrap()
959 .1
960 .datum_at(1)
961 .to_owned_datum(),
962 Some(ScalarImpl::Decimal(_))
963 ));
964 }
965
966 #[test]
967 fn test_prepare_dml_batch_payload_rejects_invalid_decoder_headers() {
968 for (header, value, expected) in [
969 (
970 "x-rw-webhook-json-timestamp-handling-mode",
971 "invalid",
972 "unrecognized value `invalid` for `x-rw-webhook-json-timestamp-handling-mode`",
973 ),
974 (
975 "x-rw-webhook-json-timestamptz-handling-mode",
976 "invalid",
977 "invalid webhook JSON decoder option",
978 ),
979 (
980 "x-rw-webhook-json-time-handling-mode",
981 "invalid",
982 "unrecognized value `invalid` for `x-rw-webhook-json-time-handling-mode`",
983 ),
984 (
985 "x-rw-webhook-json-bigint-unsigned-handling-mode",
986 "invalid",
987 "unrecognized value `invalid` for `x-rw-webhook-json-bigint-unsigned-handling-mode`",
988 ),
989 (
990 "x-rw-webhook-json-handle-toast-columns",
991 "invalid",
992 "unrecognized value `invalid` for `x-rw-webhook-json-handle-toast-columns`",
993 ),
994 ] {
995 let mut headers = HeaderMap::new();
996 headers.insert(header, HeaderValue::from_static(value));
997 let raw_dml_batch = raw_batch(
998 21,
999 vec![raw_req(Some("upsert"), r#"{"id":1,"name":"alice"}"#)],
1000 );
1001 let payload_schema = PayloadSchema::FullSchema {
1002 columns: test_columns(&[
1003 ("id", DataType::Int32, true),
1004 ("name", DataType::Varchar, false),
1005 ]),
1006 };
1007
1008 let err =
1009 prepare_dml_batch_payload(&headers, raw_dml_batch, &payload_schema).unwrap_err();
1010 assert!(err.contains(expected), "{header}");
1011 }
1012 }
1013
1014 #[test]
1015 fn test_prepare_dml_batch_payload_returns_first_type_error_with_item_index() {
1016 let raw_dml_batch = raw_batch(
1017 31,
1018 vec![
1019 raw_req(Some("upsert"), r#"{"id":1,"name":"alice"}"#),
1020 raw_req(Some("upsert"), r#"{"id":"not-an-int","name":"bob"}"#),
1021 ],
1022 );
1023 let payload_schema = PayloadSchema::FullSchema {
1024 columns: test_columns(&[
1025 ("id", DataType::Int32, true),
1026 ("name", DataType::Varchar, false),
1027 ]),
1028 };
1029
1030 let err = prepare_dml_batch_payload(&HeaderMap::new(), raw_dml_batch, &payload_schema)
1031 .unwrap_err();
1032
1033 assert!(err.contains("dml_batch_id 31 item 2"));
1034 assert!(err.contains("failed to decode webhook JSON payload"));
1035 }
1036
1037 #[test]
1038 fn test_prepare_dml_batch_payload_single_jsonb_accepts_scalar_json() {
1039 let raw_dml_batch = raw_batch(41, vec![raw_req(Some("upsert"), "123")]);
1040
1041 let prepared_batch = prepare_dml_batch_payload(
1042 &HeaderMap::new(),
1043 raw_dml_batch,
1044 &PayloadSchema::SingleJsonb,
1045 )
1046 .unwrap();
1047 let payload = prepared_batch.payload;
1048 let chunk = StreamChunk::from_protobuf(payload.chunk.as_ref().unwrap()).unwrap();
1049
1050 assert_eq!(prepared_batch.dml_batch_id, 41);
1051 assert_eq!(payload.dml_batch_id, 41);
1052 assert_eq!(chunk.ops(), &[Op::Insert]);
1053 assert!(matches!(
1054 chunk.rows().next().unwrap().1.datum_at(0).to_owned_datum(),
1055 Some(ScalarImpl::Jsonb(_))
1056 ));
1057 }
1058}