use std::collections::HashMap;
use std::sync::atomic::{AtomicU64, Ordering};
use std::sync::Arc;
use serde_json::Value;
use crate::error::{CdpError, Result};
use crate::transport::{CdpEvent, Transport, TransportKind};
#[derive(Debug, Clone)]
pub struct ConnectionConfig {
pub default_timeout_ms: u64,
pub transport_kind: TransportKind,
}
impl Default for ConnectionConfig {
fn default() -> Self {
Self {
default_timeout_ms: 30_000,
transport_kind: TransportKind::InMemory,
}
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct ParsedConnectUrl {
pub raw: String,
pub scheme: String,
pub transport_kind: TransportKind,
}
impl ParsedConnectUrl {
pub fn new(raw: impl Into<String>, scheme: impl Into<String>, kind: TransportKind) -> Self {
Self {
raw: raw.into(),
scheme: scheme.into(),
transport_kind: kind,
}
}
}
pub type EventListenerId = u64;
pub type EventListener = Arc<dyn Fn(CdpEvent) + Send + Sync>;
#[derive(Clone)]
struct EventListenerEntry {
id: EventListenerId,
handler: EventListener,
once: bool,
}
static NEXT_GLOBAL_ID: AtomicU64 = AtomicU64::new(1);
fn next_command_id() -> u64 {
NEXT_GLOBAL_ID.fetch_add(1, Ordering::Relaxed)
}
pub struct Connection {
transport: Box<dyn Transport>,
config: ConnectionConfig,
event_handlers: HashMap<String, Vec<EventListenerEntry>>,
next_listener_id: EventListenerId,
closed: bool,
}
impl std::fmt::Debug for Connection {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("Connection")
.field("transport_kind", &self.transport.kind())
.field("config", &self.config)
.field("closed", &self.closed)
.field("handler_count", &self.event_handlers.len())
.finish()
}
}
impl Connection {
pub fn new(transport: Box<dyn Transport>, config: ConnectionConfig) -> Self {
Self {
transport,
config,
event_handlers: HashMap::new(),
next_listener_id: 1,
closed: false,
}
}
pub fn from_transport(transport: Box<dyn Transport>) -> Self {
let kind = transport.kind();
Self::new(
transport,
ConnectionConfig {
default_timeout_ms: 30_000,
transport_kind: kind,
},
)
}
pub fn config(&self) -> &ConnectionConfig {
&self.config
}
pub fn transport_kind(&self) -> TransportKind {
self.transport.kind()
}
pub fn is_closed(&self) -> bool {
self.closed
}
pub fn send_command(&mut self, method: &str, params: Value) -> Result<Value> {
self.send_command_with_session(method, params, None)
}
pub fn send_command_with_session(
&mut self,
method: &str,
params: Value,
session_id: Option<&str>,
) -> Result<Value> {
if self.closed {
return Err(CdpError::ConnectionClosed);
}
let _id = next_command_id(); self.transport.send_command(method, params, session_id)
}
pub fn recv_event(&mut self) -> Result<Option<CdpEvent>> {
if self.closed {
return Err(CdpError::ConnectionClosed);
}
let event = self.transport.recv_event()?;
if let Some(ref ev) = event {
self.dispatch_event(ev.clone());
}
Ok(event)
}
pub fn drain_events(&mut self) -> Result<Vec<CdpEvent>> {
let mut events = Vec::new();
loop {
match self.recv_event()? {
Some(ev) => events.push(ev),
None => break,
}
}
Ok(events)
}
pub fn on_event(&mut self, method: &str, handler: EventListener) -> EventListenerId {
let id = self.next_listener_id;
self.next_listener_id += 1;
self.event_handlers
.entry(method.to_string())
.or_default()
.push(EventListenerEntry {
id,
handler,
once: false,
});
id
}
pub fn once_event(&mut self, method: &str, handler: EventListener) -> EventListenerId {
let id = self.next_listener_id;
self.next_listener_id += 1;
self.event_handlers
.entry(method.to_string())
.or_default()
.push(EventListenerEntry {
id,
handler,
once: true,
});
id
}
pub fn off_event(&mut self, method: &str, listener_id: EventListenerId) -> bool {
let removed = if let Some(handlers) = self.event_handlers.get_mut(method) {
let before = handlers.len();
handlers.retain(|e| e.id != listener_id);
before > handlers.len()
} else {
false
};
if removed
&& self
.event_handlers
.get(method)
.map_or(false, |v| v.is_empty())
{
self.event_handlers.remove(method);
}
removed
}
pub fn remove_all_event_handlers(&mut self, method: &str) {
self.event_handlers.remove(method);
}
pub fn event_handler_count(&self, method: &str) -> usize {
self.event_handlers
.get(method)
.map(|v| v.len())
.unwrap_or(0)
}
fn dispatch_event(&mut self, event: CdpEvent) {
let to_call: Vec<(EventListenerId, EventListener, bool)> = self
.event_handlers
.get(&event.method)
.map(|v| {
v.iter()
.map(|e| (e.id, e.handler.clone(), e.once))
.collect()
})
.unwrap_or_default();
let wildcard_to_call: Vec<(EventListenerId, EventListener, bool)> = self
.event_handlers
.get("*")
.map(|v| {
v.iter()
.map(|e| (e.id, e.handler.clone(), e.once))
.collect()
})
.unwrap_or_default();
let once_ids: Vec<EventListenerId> = to_call
.iter()
.chain(wildcard_to_call.iter())
.filter(|(_, _, o)| *o)
.map(|(id, _, _)| *id)
.collect();
for (_, handler, _) in &to_call {
handler(event.clone());
}
for (_, handler, _) in &wildcard_to_call {
handler(event.clone());
}
if !once_ids.is_empty() {
for method in [&event.method[..], "*"] {
if let Some(list) = self.event_handlers.get_mut(method) {
list.retain(|e| !once_ids.contains(&e.id));
if list.is_empty() {
self.event_handlers.remove(method);
}
}
}
}
}
pub fn close(&mut self) -> Result<()> {
if !self.closed {
self.closed = true;
self.transport.close()?;
self.event_handlers.clear();
}
Ok(())
}
pub fn transport_mut(&mut self) -> &mut dyn Transport {
&mut *self.transport
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::sync::atomic::AtomicU32;
struct MockTransport {
kind: TransportKind,
closed: bool,
next_response: Option<Value>,
event_queue: std::collections::VecDeque<CdpEvent>,
}
impl Transport for MockTransport {
fn kind(&self) -> TransportKind {
self.kind
}
fn send_command(
&mut self,
_method: &str,
_params: Value,
_session_id: Option<&str>,
) -> Result<Value> {
if self.closed {
return Err(CdpError::ConnectionClosed);
}
self.next_response
.clone()
.ok_or(CdpError::ProtocolError("no response".into()))
}
fn recv_event(&mut self) -> Result<Option<CdpEvent>> {
if self.closed {
return Err(CdpError::ConnectionClosed);
}
Ok(self.event_queue.pop_front())
}
fn close(&mut self) -> Result<()> {
self.closed = true;
Ok(())
}
}
fn make_mock_with_response(response: Value) -> Connection {
let mock = MockTransport {
kind: TransportKind::InMemory,
closed: false,
next_response: Some(response),
event_queue: std::collections::VecDeque::new(),
};
Connection::from_transport(Box::new(mock))
}
fn make_mock_with_events(events: Vec<CdpEvent>) -> Connection {
let mock = MockTransport {
kind: TransportKind::InMemory,
closed: false,
next_response: Some(Value::Null),
event_queue: events.into_iter().collect(),
};
Connection::from_transport(Box::new(mock))
}
#[test]
fn connection_config_default() {
let cfg = ConnectionConfig::default();
assert_eq!(cfg.default_timeout_ms, 30_000);
assert_eq!(cfg.transport_kind, TransportKind::InMemory);
}
#[test]
fn parsed_connect_url_construction() {
let parsed = ParsedConnectUrl::new("memory://bao", "memory", TransportKind::InMemory);
assert_eq!(parsed.raw, "memory://bao");
assert_eq!(parsed.scheme, "memory");
assert_eq!(parsed.transport_kind, TransportKind::InMemory);
}
#[test]
fn connection_new_carries_config() {
let cfg = ConnectionConfig {
default_timeout_ms: 5000,
transport_kind: TransportKind::WebSocket,
};
let mock = MockTransport {
kind: TransportKind::WebSocket,
closed: false,
next_response: Some(Value::Null),
event_queue: std::collections::VecDeque::new(),
};
let conn = Connection::new(Box::new(mock), cfg);
assert_eq!(conn.config().default_timeout_ms, 5000);
assert_eq!(conn.config().transport_kind, TransportKind::WebSocket);
assert_eq!(conn.transport_kind(), TransportKind::WebSocket);
}
#[test]
fn connection_send_command_returns_response() {
let mut conn = make_mock_with_response(serde_json::json!({"url": "https://example.com"}));
let result = conn
.send_command("Page.navigate", serde_json::json!({}))
.unwrap();
assert_eq!(result["url"], "https://example.com");
}
#[test]
fn connection_send_command_after_close_returns_error() {
let mut conn = make_mock_with_response(Value::Null);
conn.close().unwrap();
let err = conn.send_command("X.y", serde_json::json!({})).unwrap_err();
assert!(matches!(err, CdpError::ConnectionClosed));
}
#[test]
fn connection_recv_event_returns_event() {
let ev = CdpEvent::new("Page.frameNavigated", serde_json::json!({"url": "x"}));
let mut conn = make_mock_with_events(vec![ev]);
let got = conn.recv_event().unwrap().expect("expected event");
assert_eq!(got.method, "Page.frameNavigated");
}
#[test]
fn connection_recv_event_none_on_empty() {
let mut conn = make_mock_with_response(Value::Null);
let got = conn.recv_event().unwrap();
assert!(got.is_none());
}
#[test]
fn connection_recv_event_after_close_returns_error() {
let mut conn = make_mock_with_response(Value::Null);
conn.close().unwrap();
let err = conn.recv_event().unwrap_err();
assert!(matches!(err, CdpError::ConnectionClosed));
}
#[test]
fn connection_event_handler_registration_and_dispatch() {
let ev = CdpEvent::new("Page.load", serde_json::json!({"url": "x"}));
let mut conn = make_mock_with_events(vec![ev]);
let counter = Arc::new(AtomicU32::new(0));
let c = counter.clone();
let handler: EventListener = Arc::new(move |_ev| {
c.fetch_add(1, Ordering::SeqCst);
});
conn.on_event("Page.load", handler);
let got = conn.recv_event().unwrap().expect("expected event");
assert_eq!(got.method, "Page.load");
assert_eq!(counter.load(Ordering::SeqCst), 1);
}
#[test]
fn connection_event_handler_off() {
let mut conn = make_mock_with_response(Value::Null);
let handler: EventListener = Arc::new(|_| {});
let id = conn.on_event("X.y", handler);
assert_eq!(conn.event_handler_count("X.y"), 1);
let removed = conn.off_event("X.y", id);
assert!(removed);
assert_eq!(conn.event_handler_count("X.y"), 0);
}
#[test]
fn connection_close_is_idempotent() {
let mut conn = make_mock_with_response(Value::Null);
conn.close().unwrap();
conn.close().unwrap();
assert!(conn.is_closed());
}
#[test]
fn connection_drain_events() {
let events = vec![
CdpEvent::new("A", Value::Null),
CdpEvent::new("B", Value::Null),
];
let mut conn = make_mock_with_events(events);
let drained = conn.drain_events().unwrap();
assert_eq!(drained.len(), 2);
assert_eq!(drained[0].method, "A");
assert_eq!(drained[1].method, "B");
}
#[test]
fn connection_wildcard_handler() {
let ev = CdpEvent::new("Page.load", Value::Null);
let mut conn = make_mock_with_events(vec![ev]);
let counter = Arc::new(AtomicU32::new(0));
let c = counter.clone();
let handler: EventListener = Arc::new(move |_ev| {
c.fetch_add(1, Ordering::SeqCst);
});
conn.on_event("*", handler);
let _ = conn.recv_event().unwrap();
assert_eq!(counter.load(Ordering::SeqCst), 1);
}
#[test]
fn next_command_id_monotonic() {
let id1 = next_command_id();
let id2 = next_command_id();
assert!(id2 > id1);
}
#[test]
fn connection_once_event_auto_removes_after_first_invocation() {
let ev1 = CdpEvent::new("Page.load", serde_json::json!({"url": "a"}));
let ev2 = CdpEvent::new("Page.load", serde_json::json!({"url": "b"}));
let mut conn = make_mock_with_events(vec![ev1, ev2]);
let counter = Arc::new(AtomicU32::new(0));
let c = counter.clone();
let handler: EventListener = Arc::new(move |_ev| {
c.fetch_add(1, Ordering::SeqCst);
});
conn.once_event("Page.load", handler);
assert_eq!(conn.event_handler_count("Page.load"), 1);
let got = conn.recv_event().unwrap().expect("expected event");
assert_eq!(got.method, "Page.load");
assert_eq!(counter.load(Ordering::SeqCst), 1);
assert_eq!(conn.event_handler_count("Page.load"), 0);
let got2 = conn.recv_event().unwrap().expect("expected event");
assert_eq!(got2.method, "Page.load");
assert_eq!(counter.load(Ordering::SeqCst), 1);
}
#[test]
fn connection_once_event_wildcard_auto_removes() {
let ev = CdpEvent::new("Page.load", Value::Null);
let mut conn = make_mock_with_events(vec![ev]);
let counter = Arc::new(AtomicU32::new(0));
let c = counter.clone();
let handler: EventListener = Arc::new(move |_ev| {
c.fetch_add(1, Ordering::SeqCst);
});
conn.once_event("*", handler);
assert_eq!(conn.event_handler_count("*"), 1);
let _ = conn.recv_event().unwrap();
assert_eq!(counter.load(Ordering::SeqCst), 1);
assert_eq!(conn.event_handler_count("*"), 0);
}
#[test]
fn connection_on_event_persists_across_events() {
let ev1 = CdpEvent::new("Page.load", Value::Null);
let ev2 = CdpEvent::new("Page.load", Value::Null);
let mut conn = make_mock_with_events(vec![ev1, ev2]);
let counter = Arc::new(AtomicU32::new(0));
let c = counter.clone();
let handler: EventListener = Arc::new(move |_ev| {
c.fetch_add(1, Ordering::SeqCst);
});
conn.on_event("Page.load", handler);
let _ = conn.recv_event().unwrap();
assert_eq!(counter.load(Ordering::SeqCst), 1);
let _ = conn.recv_event().unwrap();
assert_eq!(counter.load(Ordering::SeqCst), 2);
assert_eq!(conn.event_handler_count("Page.load"), 1);
}
}