Skip to main content

laminar_connectors/websocket/
mod.rs

1//! WebSocket client source plus client/server sinks.
2//!
3//! WebSocket is non-replayable. Sources and sinks are best-effort, and failures
4//! can produce gaps or duplicates.
5
6mod backpressure;
7mod connection;
8mod fanout;
9mod metrics;
10mod parser;
11mod protocol;
12mod serializer;
13mod sink;
14mod sink_client;
15mod sink_config;
16mod sink_metrics;
17mod source;
18mod source_config;
19
20use std::sync::{Arc, OnceLock, Weak};
21
22use metrics::WebSocketSourceMetrics;
23use sink::WebSocketSinkServer;
24use sink_client::WebSocketSinkClient;
25use sink_config::WebSocketSinkConfig;
26use sink_metrics::WebSocketSinkMetrics;
27use source::WebSocketSource;
28use source_config::WebSocketSourceConfig;
29
30use crate::config::{ConfigKeySpec, ConnectorInfo};
31use crate::registry::ConnectorRegistry;
32
33struct RegisteredMetricFamily<T> {
34    registry: Weak<prometheus::Registry>,
35    metrics: T,
36}
37
38struct MetricFamilyCache<T> {
39    family: OnceLock<RegisteredMetricFamily<T>>,
40    initialization: parking_lot::Mutex<()>,
41}
42
43impl<T> Default for MetricFamilyCache<T> {
44    fn default() -> Self {
45        Self {
46            family: OnceLock::new(),
47            initialization: parking_lot::Mutex::new(()),
48        }
49    }
50}
51
52impl<T: Clone> MetricFamilyCache<T> {
53    fn get_or_try_init(
54        &self,
55        registry: &Arc<prometheus::Registry>,
56        initialize: impl FnOnce() -> Result<T, crate::error::ConnectorError>,
57    ) -> Result<T, crate::error::ConnectorError> {
58        if let Some(family) = self.family.get() {
59            return family.for_registry(registry);
60        }
61        let _guard = self.initialization.lock();
62        if let Some(family) = self.family.get() {
63            return family.for_registry(registry);
64        }
65        let family = RegisteredMetricFamily {
66            registry: Arc::downgrade(registry),
67            metrics: initialize()?,
68        };
69        self.family
70            .set(family)
71            .unwrap_or_else(|_| unreachable!("metric initialization is serialized"));
72        Ok(self
73            .family
74            .get()
75            .expect("metric family was initialized")
76            .metrics
77            .clone())
78    }
79}
80
81impl<T: Clone> RegisteredMetricFamily<T> {
82    fn for_registry(
83        &self,
84        registry: &Arc<prometheus::Registry>,
85    ) -> Result<T, crate::error::ConnectorError> {
86        if !Weak::ptr_eq(&self.registry, &Arc::downgrade(registry)) {
87            return Err(crate::error::ConnectorError::ConfigurationError(
88                "WebSocket connector registry is already bound to a different Prometheus registry"
89                    .into(),
90            ));
91        }
92        Ok(self.metrics.clone())
93    }
94}
95
96/// Registers the WebSocket source connector with the given registry.
97///
98/// After registration, the runtime can instantiate `WebSocketSource` by
99/// name when processing `CREATE SOURCE ... WITH (connector = 'websocket')`.
100///
101/// # Errors
102///
103/// Returns an error when the source name is already registered or the
104/// connector registry is frozen.
105pub fn register_websocket_source(
106    registry: &ConnectorRegistry,
107) -> Result<(), crate::error::ConnectorError> {
108    let info = ConnectorInfo {
109        name: "websocket".to_string(),
110        display_name: "WebSocket Source".to_string(),
111        version: env!("CARGO_PKG_VERSION").to_string(),
112        is_source: true,
113        is_sink: false,
114        config_keys: websocket_source_config_keys(),
115    };
116
117    let metric_cache = Arc::new(MetricFamilyCache::<WebSocketSourceMetrics>::default());
118    registry.register_source(
119        "websocket",
120        info,
121        Arc::new(move |registry: Option<&Arc<prometheus::Registry>>| {
122            let metrics = if let Some(registry) = registry {
123                metric_cache
124                    .get_or_try_init(registry, || WebSocketSourceMetrics::register(registry))?
125            } else {
126                WebSocketSourceMetrics::local()
127            };
128            Ok(Box::new(WebSocketSource::new(
129                Arc::new(arrow_schema::Schema::empty()),
130                WebSocketSourceConfig::default(),
131                metrics,
132            )))
133        }),
134    )
135}
136
137/// Registers the WebSocket sink connector with the given registry.
138///
139/// The sink factory selects server or client mode from the validated config
140/// before either implementation performs network I/O.
141///
142/// # Errors
143///
144/// Returns an error when the sink name is already registered or the connector
145/// registry is frozen.
146pub fn register_websocket_sink(
147    registry: &ConnectorRegistry,
148) -> Result<(), crate::error::ConnectorError> {
149    let info = ConnectorInfo {
150        name: "websocket".to_string(),
151        display_name: "WebSocket Sink".to_string(),
152        version: env!("CARGO_PKG_VERSION").to_string(),
153        is_source: false,
154        is_sink: true,
155        config_keys: websocket_sink_config_keys(),
156    };
157
158    let metric_cache = Arc::new(MetricFamilyCache::<WebSocketSinkMetrics>::default());
159    registry.register_sink(
160        "websocket",
161        info,
162        Arc::new(
163            move |config, registry: Option<&Arc<prometheus::Registry>>| {
164                let sink_config = WebSocketSinkConfig::from_config(config)?;
165                let decoded_schema = config.arrow_schema();
166                if config.get("_arrow_schema").is_some() && decoded_schema.is_none() {
167                    return Err(crate::error::ConnectorError::ConfigurationError(
168                        "invalid WebSocket sink _arrow_schema encoding".into(),
169                    ));
170                }
171                let schema = decoded_schema.ok_or_else(|| {
172                    crate::error::ConnectorError::ConfigurationError(
173                        "WebSocket sink requires a declared Arrow schema".into(),
174                    )
175                })?;
176                let metrics = if let Some(registry) = registry {
177                    metric_cache
178                        .get_or_try_init(registry, || WebSocketSinkMetrics::register(registry))?
179                } else {
180                    WebSocketSinkMetrics::local()
181                };
182                let is_server = matches!(&sink_config, WebSocketSinkConfig::Server { .. });
183                let sink: Box<dyn crate::connector::SinkConnector> = if is_server {
184                    Box::new(WebSocketSinkServer::new(schema, sink_config, metrics))
185                } else {
186                    Box::new(WebSocketSinkClient::new(schema, sink_config, metrics))
187                };
188                Ok(sink)
189            },
190        ),
191    )
192}
193
194fn websocket_source_config_keys() -> Vec<ConfigKeySpec> {
195    vec![
196        ConfigKeySpec::required("url", "WebSocket URL to connect to (ws:// or wss://)"),
197        ConfigKeySpec::optional("format", "Message format (json/csv/binary)", "json"),
198        ConfigKeySpec::optional(
199            "subscribe.message",
200            "Text subscription message to send after handshake",
201            "",
202        ),
203        ConfigKeySpec::optional("reconnect.enabled", "Enable automatic reconnection", "true"),
204        ConfigKeySpec::optional(
205            "reconnect.initial.delay.ms",
206            "Initial reconnect delay in ms",
207            "100",
208        ),
209        ConfigKeySpec::optional(
210            "reconnect.max.delay.ms",
211            "Maximum reconnect delay in ms",
212            "30000",
213        ),
214        ConfigKeySpec::optional(
215            "reconnect.max.retries",
216            "Maximum reconnect attempts; empty means unlimited",
217            "",
218        ),
219        ConfigKeySpec::optional(
220            "on.backpressure",
221            "Backpressure strategy (block/drop_newest)",
222            "block",
223        ),
224        ConfigKeySpec::optional(
225            "max.message.size",
226            "Max WebSocket message size in bytes",
227            "67108864",
228        ),
229    ]
230}
231
232fn websocket_sink_config_keys() -> Vec<ConfigKeySpec> {
233    vec![
234        ConfigKeySpec::optional(
235            "bind.address",
236            "Socket address required in server mode (e.g., 0.0.0.0:8080)",
237            "",
238        ),
239        ConfigKeySpec::optional("mode", "Operating mode (server/client)", "server"),
240        ConfigKeySpec::optional(
241            "max.connections",
242            "Max concurrent client connections",
243            "10000",
244        ),
245        ConfigKeySpec::optional("ping.interval.ms", "Ping interval in ms", "30000"),
246        ConfigKeySpec::optional("ping.timeout.ms", "Pong timeout in ms", "10000"),
247        ConfigKeySpec::optional("url", "WebSocket URL required in client mode", ""),
248    ]
249}
250
251#[cfg(test)]
252mod tests {
253    use super::*;
254    use arrow_schema::{DataType, Field, Schema};
255    use std::thread;
256
257    fn config_with_schema(mode: &str) -> crate::config::ConnectorConfig {
258        let mut config = crate::config::ConnectorConfig::new("websocket");
259        config.set("mode", mode);
260        let schema = Schema::new(vec![Field::new("payload", DataType::Utf8, false)]);
261        config.set(
262            "_arrow_schema",
263            crate::config::encode_arrow_schema_ipc(&schema),
264        );
265        config
266    }
267
268    fn factory_error(
269        registry: &ConnectorRegistry,
270        config: &crate::config::ConnectorConfig,
271    ) -> String {
272        match registry.create_sink(config, None) {
273            Ok(_) => panic!("expected sink factory error"),
274            Err(error) => error.to_string(),
275        }
276    }
277
278    #[test]
279    fn sink_factory_dispatches_server_without_opening_a_socket() {
280        let registry = ConnectorRegistry::new();
281        register_websocket_sink(&registry).unwrap();
282        let mut config = config_with_schema("server");
283        config.set("bind.address", "127.0.0.1:0");
284
285        let sink = registry.create_sink(&config, None).unwrap();
286
287        assert!(sink.contract(&config).is_ok());
288        assert_eq!(sink.schema().field(0).name(), "payload");
289    }
290
291    #[test]
292    fn sink_factory_dispatches_client_without_opening_a_socket() {
293        let registry = ConnectorRegistry::new();
294        register_websocket_sink(&registry).unwrap();
295        let mut config = config_with_schema("client");
296        config.set("url", "wss://example.test/events");
297
298        let sink = registry.create_sink(&config, None).unwrap();
299
300        assert!(sink.contract(&config).is_ok());
301        assert_eq!(sink.schema().field(0).name(), "payload");
302    }
303
304    #[test]
305    fn sink_factory_rejects_missing_or_invalid_mode_specific_config() {
306        let registry = ConnectorRegistry::new();
307        register_websocket_sink(&registry).unwrap();
308
309        let missing_server_bind = config_with_schema("server");
310        assert!(factory_error(&registry, &missing_server_bind).contains("bind.address"));
311
312        let mut invalid_server_bind = config_with_schema("server");
313        invalid_server_bind.set("bind.address", "not-a-socket-address");
314        assert!(factory_error(&registry, &invalid_server_bind).contains("invalid WebSocket server"));
315
316        let mut missing_client_url = config_with_schema("client");
317        assert!(factory_error(&registry, &missing_client_url).contains("url"));
318
319        missing_client_url.set("url", "https://example.test/not-websocket");
320        assert!(factory_error(&registry, &missing_client_url).contains("expected ws:// or wss://"));
321
322        let mut invalid_schema = config_with_schema("server");
323        invalid_schema.set("bind.address", "127.0.0.1:0");
324        invalid_schema.set("_arrow_schema", "not-arrow-ipc");
325        assert!(factory_error(&registry, &invalid_schema).contains("_arrow_schema"));
326
327        let mut missing_schema = crate::config::ConnectorConfig::new("websocket");
328        missing_schema.set("mode", "server");
329        missing_schema.set("bind.address", "127.0.0.1:0");
330        assert!(factory_error(&registry, &missing_schema).contains("declared Arrow schema"));
331    }
332
333    #[test]
334    fn bind_address_metadata_is_mode_conditional() {
335        let registry = ConnectorRegistry::new();
336        register_websocket_sink(&registry).unwrap();
337        let info = registry.sink_info("websocket").unwrap();
338        let bind = info
339            .config_keys
340            .iter()
341            .find(|key| key.key == "bind.address")
342            .unwrap();
343        assert!(!bind.required);
344    }
345
346    #[test]
347    fn source_metadata_exposes_only_runtime_options() {
348        let registry = ConnectorRegistry::new();
349        register_websocket_source(&registry).unwrap();
350        let info = registry.source_info("websocket").unwrap();
351        let keys: std::collections::HashSet<&str> = info
352            .config_keys
353            .iter()
354            .map(|key| key.key.as_str())
355            .collect();
356        let expected: std::collections::HashSet<&str> = [
357            "url",
358            "format",
359            "subscribe.message",
360            "reconnect.enabled",
361            "reconnect.initial.delay.ms",
362            "reconnect.max.delay.ms",
363            "reconnect.max.retries",
364            "on.backpressure",
365            "max.message.size",
366        ]
367        .into_iter()
368        .collect();
369        assert_eq!(keys, expected);
370    }
371
372    #[test]
373    fn sink_metadata_exposes_only_runtime_options() {
374        let registry = ConnectorRegistry::new();
375        register_websocket_sink(&registry).unwrap();
376        let info = registry.sink_info("websocket").unwrap();
377        let keys: std::collections::HashSet<&str> = info
378            .config_keys
379            .iter()
380            .map(|key| key.key.as_str())
381            .collect();
382        let expected: std::collections::HashSet<&str> = [
383            "bind.address",
384            "mode",
385            "max.connections",
386            "ping.interval.ms",
387            "ping.timeout.ms",
388            "url",
389        ]
390        .into_iter()
391        .collect();
392        assert_eq!(keys, expected);
393    }
394
395    #[test]
396    fn source_factory_registers_one_shared_metric_family() {
397        let connectors = ConnectorRegistry::new();
398        register_websocket_source(&connectors).unwrap();
399        let metrics = Arc::new(prometheus::Registry::new());
400        let config = crate::config::ConnectorConfig::new("websocket");
401
402        connectors.create_source(&config, Some(&metrics)).unwrap();
403        connectors.create_source(&config, Some(&metrics)).unwrap();
404
405        let families = metrics.gather();
406        assert_eq!(
407            families
408                .iter()
409                .filter(|family| family.name().starts_with("websocket_source_"))
410                .count(),
411            5
412        );
413    }
414
415    #[test]
416    fn metric_family_cache_is_concurrent_and_shared() {
417        let cache = Arc::new(MetricFamilyCache::<WebSocketSourceMetrics>::default());
418        let registry = Arc::new(prometheus::Registry::new());
419        let workers = (0..16)
420            .map(|_| {
421                let cache = Arc::clone(&cache);
422                let registry = Arc::clone(&registry);
423                thread::spawn(move || {
424                    let metrics = cache
425                        .get_or_try_init(&registry, || WebSocketSourceMetrics::register(&registry))
426                        .unwrap();
427                    metrics.record_message(1);
428                })
429            })
430            .collect::<Vec<_>>();
431
432        for worker in workers {
433            worker.join().unwrap();
434        }
435
436        assert_eq!(
437            cache.family.get().unwrap().metrics.messages_received.get(),
438            16
439        );
440    }
441
442    #[test]
443    fn source_factory_rejects_a_different_metrics_registry() {
444        let connectors = ConnectorRegistry::new();
445        register_websocket_source(&connectors).unwrap();
446        let first = Arc::new(prometheus::Registry::new());
447        let second = Arc::new(prometheus::Registry::new());
448        let config = crate::config::ConnectorConfig::new("websocket");
449
450        connectors.create_source(&config, Some(&first)).unwrap();
451        let error = match connectors.create_source(&config, Some(&second)) {
452            Ok(_) => panic!("expected metrics registry mismatch"),
453            Err(error) => error.to_string(),
454        };
455
456        assert!(error.contains("different Prometheus registry"), "{error}");
457    }
458}