Skip to main content

gthings_cdp/
connection.rs

1use futures_util::{SinkExt, StreamExt};
2use serde_json::Value;
3use std::collections::HashMap;
4use std::sync::Arc;
5use std::sync::atomic::{AtomicU64, Ordering};
6use std::time::Duration;
7use tokio::net::TcpStream;
8use tokio::sync::{Mutex, oneshot};
9use tokio_tungstenite::MaybeTlsStream;
10use tokio_tungstenite::WebSocketStream;
11use tokio_tungstenite::tungstenite::Message;
12use tracing;
13
14use crate::error::{CdpError, Result};
15
16type PendingMap = Arc<Mutex<HashMap<u64, oneshot::Sender<Value>>>>;
17
18/// CDP WebSocket connection dispatching commands via oneshot channels.
19pub struct Connection {
20    write: futures_util::stream::SplitSink<WebSocketStream<MaybeTlsStream<TcpStream>>, Message>,
21    pending: PendingMap,
22    next_id: AtomicU64,
23}
24
25impl Connection {
26    /// Create a new CDP connection. Spawns a background reader for response dispatch.
27    pub async fn new(
28        ws_stream: WebSocketStream<MaybeTlsStream<TcpStream>>,
29        mut kill_rx: tokio::sync::oneshot::Receiver<()>,
30    ) -> Result<Self> {
31        let (write, read) = ws_stream.split();
32        let pending: PendingMap = Arc::new(Mutex::new(HashMap::new()));
33        let pending_clone = pending.clone();
34
35        tokio::spawn(async move {
36            let mut read = read;
37            loop {
38                tokio::select! {
39                    msg = read.next() => {
40                        match msg {
41                            Some(Ok(Message::Text(text))) => {
42                                if let Ok(value) = serde_json::from_str::<Value>(&text) {
43                                    if let Some(id) = value.get("id").and_then(|v| v.as_u64()) {
44                                        let mut map = pending_clone.lock().await;
45                                        if let Some(tx) = map.remove(&id) {
46                                            let _ = tx.send(value);
47                                        }
48                                    }
49                                    // Events (no "id" field) are silently ignored
50                                }
51                            }
52                            Some(Ok(Message::Binary(_))) => {
53                                // Binary frames not expected from CDP
54                            }
55                            Some(Ok(Message::Close(frame))) => {
56                                tracing::debug!("CDP WebSocket closed: {:?}", frame);
57                                break;
58                            }
59                            Some(Err(e)) => {
60                                tracing::warn!("CDP WebSocket error: {e}");
61                                break;
62                            }
63                            None => {
64                                tracing::debug!("CDP WebSocket stream ended");
65                                break;
66                            }
67                            _ => {}
68                        }
69                    }
70                    _ = &mut kill_rx => {
71                        tracing::debug!("Kill signal received, stopping reader");
72                        break;
73                    }
74                }
75            }
76        });
77
78        Ok(Connection {
79            write,
80            pending,
81            next_id: AtomicU64::new(1),
82        })
83    }
84
85    /// Send a CDP command and wait for the response via oneshot dispatch.
86    pub async fn call(&mut self, method: &str, params: Value) -> Result<Value> {
87        let id = self.next_id.fetch_add(1, Ordering::Relaxed);
88
89        let cmd = serde_json::json!({
90            "id": id,
91            "method": method,
92            "params": params,
93        });
94
95        let text = serde_json::to_string(&cmd)?;
96
97        let (tx, rx) = oneshot::channel();
98        {
99            let mut map = self.pending.lock().await;
100            map.insert(id, tx);
101        }
102
103        self.write.send(Message::Text(text)).await?;
104
105        let response = tokio::time::timeout(Duration::from_secs(30), rx)
106            .await
107            .map_err(|_| CdpError::Timeout(30000))?
108            .map_err(|_| CdpError::ChannelBroken)?;
109
110        if let Some(err) = response.get("error") {
111            let msg = err
112                .get("message")
113                .and_then(|v| v.as_str())
114                .unwrap_or("unknown error")
115                .to_string();
116            return Err(CdpError::CommandFailed {
117                method: method.to_string(),
118                msg,
119            });
120        }
121
122        Ok(response.get("result").cloned().unwrap_or(Value::Null))
123    }
124
125    /// Send a CDP command with an explicit sessionId (for tab-specific commands).
126    pub async fn call_with_session(
127        &mut self,
128        session_id: &str,
129        method: &str,
130        params: Value,
131    ) -> Result<Value> {
132        let id = self.next_id.fetch_add(1, Ordering::Relaxed);
133
134        let cmd = serde_json::json!({
135            "id": id,
136            "method": method,
137            "params": params,
138            "sessionId": session_id,
139        });
140
141        let text = serde_json::to_string(&cmd)?;
142
143        let (tx, rx) = oneshot::channel();
144        {
145            let mut map = self.pending.lock().await;
146            map.insert(id, tx);
147        }
148
149        self.write.send(Message::Text(text)).await?;
150
151        let response = tokio::time::timeout(Duration::from_secs(30), rx)
152            .await
153            .map_err(|_| CdpError::Timeout(30000))?
154            .map_err(|_| CdpError::ChannelBroken)?;
155
156        if let Some(err) = response.get("error") {
157            let msg = err
158                .get("message")
159                .and_then(|v| v.as_str())
160                .unwrap_or("unknown error")
161                .to_string();
162            return Err(CdpError::CommandFailed {
163                method: method.to_string(),
164                msg,
165            });
166        }
167
168        Ok(response.get("result").cloned().unwrap_or(Value::Null))
169    }
170}
171
172#[cfg(test)]
173mod tests {
174    use serde_json::json;
175
176    #[test]
177    fn test_cdp_response_parsing() {
178        let response = json!({"id": 1, "result": {"value": "hello"}});
179        assert_eq!(response["id"].as_u64(), Some(1));
180        assert!(response.get("error").is_none());
181        assert_eq!(response["result"]["value"].as_str(), Some("hello"));
182    }
183
184    #[test]
185    fn test_cdp_error_response() {
186        let response =
187            json!({"id": 2, "error": {"code": -32000, "message": "Cannot find context"}});
188        assert!(response.get("error").is_some());
189        assert_eq!(
190            response["error"]["message"].as_str(),
191            Some("Cannot find context")
192        );
193    }
194
195    #[test]
196    fn test_cdp_event_has_no_id() {
197        let event = json!({"method": "Page.frameStartedLoading", "params": {"frameId": "123"}});
198        assert!(event.get("id").is_none());
199        assert!(event.get("method").is_some());
200        assert!(event.get("params").is_some());
201    }
202}