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