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("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
169// Keep licensing at the WebSocket entry point so webhook ingestion remains unlicensed.
170fn 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}