Skip to main content

risingwave_frontend/webhook/
websocket.rs

1// Copyright 2026 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
15//! WebSocket-based streaming DML ingest endpoint.
16//!
17//! Clients open a `WebSocket` connection per table. The first text frame authenticates the session
18//! and later frames carry unsigned batches of upsert / delete DML messages. Each websocket
19//! payload carries a monotonically increasing `dml_batch_id`.
20//!
21//! Wire format (JSON over text `WebSocket` frames):
22//!
23//! **Client → Server** (authenticated init):
24//! ```json
25//! {"type": "init", "timestamp": 1760000000000}
26//! ```
27//!
28//! **Client → Server** (DML batch):
29//! ```json
30//! {
31//!   "dml_batch_id": 1,
32//!   "items": [
33//!     {"op": "upsert", "data": {"id": 1, "name": "foo"}},
34//!     {"op": "delete", "data": {"id": 1, "name": "foo"}}
35//!   ]
36//! }
37//! ```
38//!
39//! **Server → Client** (ack):
40//! ```json
41//! {"ack": 1}
42//! ```
43//!
44//! **Server → Client** (fatal — the connection will close):
45//! ```json
46//! {"fatal": "DML channel closed, please reconnect"}
47//! ```
48use 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
87/// Shared state for the ingest service.
88pub 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    // `RawValue` is unsized; `Box<RawValue>` is serde's owned form that preserves the original JSON text.
119    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
173// Keep licensing at the WebSocket entry point so webhook ingestion remains unlicensed.
174fn 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}