1mod 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
96pub 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
137pub 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(®istry).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(®istry).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(®istry).unwrap();
308
309 let missing_server_bind = config_with_schema("server");
310 assert!(factory_error(®istry, &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(®istry, &invalid_server_bind).contains("invalid WebSocket server"));
315
316 let mut missing_client_url = config_with_schema("client");
317 assert!(factory_error(®istry, &missing_client_url).contains("url"));
318
319 missing_client_url.set("url", "https://example.test/not-websocket");
320 assert!(factory_error(®istry, &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(®istry, &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(®istry, &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(®istry).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(®istry).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(®istry).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(®istry);
423 thread::spawn(move || {
424 let metrics = cache
425 .get_or_try_init(®istry, || WebSocketSourceMetrics::register(®istry))
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}