Skip to main content

agent_first_http/sdk/cdp/
ws_client.rs

1//! Async CDP WebSocket client.
2//!
3//! Each `Connection` owns one WebSocket to the host's `/cdp` endpoint.
4//! Sent commands are tagged with monotonically increasing ids; replies and
5//! events are demuxed by the reader task into pending-request channels and
6//! a broadcast event stream respectively.
7
8use std::collections::HashMap;
9use std::sync::atomic::{AtomicI64, Ordering};
10use std::sync::Arc;
11use std::time::Duration;
12
13use futures::{SinkExt, StreamExt};
14use serde::Serialize;
15use serde_json::Value;
16use tokio::io::{AsyncRead, AsyncWrite};
17use tokio::sync::{broadcast, mpsc, oneshot, Mutex};
18use tokio::task::JoinHandle;
19use tokio_tungstenite::tungstenite::{
20    self,
21    client::IntoClientRequest,
22    handshake::client::generate_key,
23    http::{header, Request, Uri},
24};
25use tokio_tungstenite::WebSocketStream;
26
27use crate::sdk::endpoint::Endpoint;
28use crate::shared::error::{Error, ErrorCode};
29
30type ReplySender = oneshot::Sender<Result<Value, CdpRemoteError>>;
31type PendingMap = Arc<Mutex<HashMap<i64, ReplySender>>>;
32
33/// One CDP connection.
34pub struct Connection {
35    tx: mpsc::UnboundedSender<OutMsg>,
36    pending: PendingMap,
37    events_tx: broadcast::Sender<CdpEvent>,
38    next_id: AtomicI64,
39    _reader: JoinHandle<()>,
40    _writer: JoinHandle<()>,
41}
42
43enum OutMsg {
44    Text(String),
45    Close,
46}
47
48#[derive(Debug, Clone)]
49pub struct CdpEvent {
50    pub method: String,
51    pub session_id: Option<String>,
52    pub params: Value,
53}
54
55#[derive(Debug, Clone)]
56pub struct CdpRemoteError {
57    pub code: i64,
58    pub message: String,
59}
60
61impl std::fmt::Display for CdpRemoteError {
62    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
63        write!(f, "CDP error {}: {}", self.code, self.message)
64    }
65}
66
67impl Connection {
68    /// Open a CDP connection from a parsed endpoint.
69    pub async fn connect_endpoint(endpoint: &Endpoint, token: Option<&str>) -> Result<Self, Error> {
70        match endpoint {
71            #[cfg(unix)]
72            Endpoint::Unix { path } => Self::connect_unix(path, token).await,
73            _ => Self::connect(&endpoint.cdp_ws_url(), token).await,
74        }
75    }
76
77    /// Open a CDP connection to `ws://endpoint/cdp` (or `wss://`),
78    /// attaching the optional bearer token via the `?token_secret=` query
79    /// parameter (the `_secret` suffix lets AFDATA redaction scrub it).
80    pub async fn connect(endpoint_ws_url: &str, token: Option<&str>) -> Result<Self, Error> {
81        // Append ?token_secret= if needed.
82        let url = match token {
83            Some(t) => append_query_pairs(endpoint_ws_url, &[("token_secret", t)])?,
84            None => endpoint_ws_url.to_string(),
85        };
86        let request = build_ws_request(&url, token)?;
87        let uri: Uri = url
88            .parse()
89            .map_err(|e| Error::new(ErrorCode::InvalidEndpoint, format!("CDP url {url:?}: {e}")))?;
90        let secure = uri
91            .scheme_str()
92            .is_some_and(|s| s.eq_ignore_ascii_case("wss"));
93        if !secure {
94            // Plaintext ws:// (all local CDP, and ws:// remote hosts): connect the
95            // TCP stream directly and run the handshake with client_async. We
96            // avoid connect_async because, with the rustls-tls-native-roots
97            // feature, it builds a TLS connector and loads the OS root-cert store
98            // even for ws:// — wasted work for every CDP connect, and stack-heavy
99            // enough to overflow Windows' 1 MiB main-thread stack.
100            let host = uri.host().ok_or_else(|| {
101                Error::new(
102                    ErrorCode::InvalidEndpoint,
103                    format!("CDP url has no host: {url:?}"),
104                )
105            })?;
106            let port = uri.port_u16().unwrap_or(80);
107            let stream = tokio::net::TcpStream::connect((host, port))
108                .await
109                .map_err(|e| {
110                    Error::new(
111                        ErrorCode::HostUnreachable,
112                        format!(
113                            "CDP connect {}: {e}",
114                            agent_first_data::redact_url_secrets(&url)
115                        ),
116                    )
117                })?;
118            let (ws, _resp) = tokio_tungstenite::client_async(request, stream)
119                .await
120                .map_err(|e| {
121                    Error::new(
122                        ErrorCode::HostUnreachable,
123                        format!(
124                            "CDP websocket {}: {e}",
125                            agent_first_data::redact_url_secrets(&url)
126                        ),
127                    )
128                })?;
129            return Ok(Self::from_ws(ws));
130        }
131        // wss:// — keep the TLS-capable connector.
132        let (ws, _resp) = tokio_tungstenite::connect_async(request)
133            .await
134            .map_err(|e| {
135                // Redact the token before it lands in the error envelope.
136                Error::new(
137                    ErrorCode::HostUnreachable,
138                    format!(
139                        "CDP connect {}: {e}",
140                        agent_first_data::redact_url_secrets(&url)
141                    ),
142                )
143            })?;
144        Ok(Self::from_ws(ws))
145    }
146
147    #[cfg(unix)]
148    async fn connect_unix(path: &std::path::Path, token: Option<&str>) -> Result<Self, Error> {
149        let url = match token {
150            Some(t) => append_query_pairs("ws://localhost/cdp", &[("token_secret", t)])?,
151            None => "ws://localhost/cdp".to_string(),
152        };
153        let request = build_ws_request(&url, token)?;
154        let stream = tokio::net::UnixStream::connect(path).await.map_err(|e| {
155            Error::new(
156                ErrorCode::HostUnreachable,
157                format!("CDP connect unix:{}: {e}", path.display()),
158            )
159        })?;
160        let (ws, _resp) = tokio_tungstenite::client_async(request, stream)
161            .await
162            .map_err(|e| {
163                Error::new(
164                    ErrorCode::HostUnreachable,
165                    format!("CDP websocket over unix:{}: {e}", path.display()),
166                )
167            })?;
168        Ok(Self::from_ws(ws))
169    }
170
171    fn from_ws<S>(ws: WebSocketStream<S>) -> Self
172    where
173        S: AsyncRead + AsyncWrite + Unpin + Send + 'static,
174    {
175        let (mut sink, mut stream) = ws.split();
176
177        let pending: PendingMap = Arc::new(Mutex::new(HashMap::new()));
178        let (events_tx, _events_rx) = broadcast::channel::<CdpEvent>(256);
179        let (tx, mut rx) = mpsc::unbounded_channel::<OutMsg>();
180
181        let pending_w = pending.clone();
182        let events_w = events_tx.clone();
183        let reader = tokio::spawn(async move {
184            while let Some(Ok(msg)) = stream.next().await {
185                match msg {
186                    tungstenite::Message::Text(t) => {
187                        if let Ok(v) = serde_json::from_str::<Value>(t.as_str()) {
188                            dispatch(v, &pending_w, &events_w).await;
189                        }
190                    }
191                    tungstenite::Message::Binary(_)
192                    | tungstenite::Message::Ping(_)
193                    | tungstenite::Message::Pong(_) => {}
194                    tungstenite::Message::Close(_) | tungstenite::Message::Frame(_) => break,
195                }
196            }
197            // On close, fail all pending requests so callers stop waiting.
198            let mut map = pending_w.lock().await;
199            for (_, sender) in map.drain() {
200                let _ = sender.send(Err(CdpRemoteError {
201                    code: -1,
202                    message: "CDP connection closed".into(),
203                }));
204            }
205        });
206
207        let writer = tokio::spawn(async move {
208            while let Some(out) = rx.recv().await {
209                let msg = match out {
210                    OutMsg::Text(t) => tungstenite::Message::Text(t.as_str().into()),
211                    OutMsg::Close => {
212                        let _ = sink.send(tungstenite::Message::Close(None)).await;
213                        break;
214                    }
215                };
216                if sink.send(msg).await.is_err() {
217                    break;
218                }
219            }
220        });
221
222        Self {
223            tx,
224            pending,
225            events_tx,
226            next_id: AtomicI64::new(1),
227            _reader: reader,
228            _writer: writer,
229        }
230    }
231
232    /// Subscribe to events for the lifetime of this connection.
233    pub fn subscribe(&self) -> broadcast::Receiver<CdpEvent> {
234        self.events_tx.subscribe()
235    }
236
237    /// Send a CDP command and await the result. `session_id` is `Some` when
238    /// the call is scoped to a flattened session (after Target.attachToTarget).
239    pub async fn send<P: Serialize>(
240        &self,
241        method: &str,
242        params: &P,
243        session_id: Option<&str>,
244    ) -> Result<Value, Error> {
245        let id = self.next_id.fetch_add(1, Ordering::SeqCst);
246        let body = match session_id {
247            Some(sid) => serde_json::json!({
248                "id": id,
249                "method": method,
250                "params": params,
251                "sessionId": sid,
252            }),
253            None => serde_json::json!({
254                "id": id,
255                "method": method,
256                "params": params,
257            }),
258        };
259        let serialized = serde_json::to_string(&body).map_err(|e| {
260            Error::new(
261                ErrorCode::InternalError,
262                format!("CDP send: serialize {method}: {e}"),
263            )
264        })?;
265        let (resp_tx, resp_rx) = oneshot::channel();
266        self.pending.lock().await.insert(id, resp_tx);
267        self.tx
268            .send(OutMsg::Text(serialized))
269            .map_err(|_| Error::new(ErrorCode::CdpUnavailable, "CDP writer closed before send"))?;
270        let value = resp_rx
271            .await
272            .map_err(|_| Error::new(ErrorCode::CdpUnavailable, "CDP reader closed"))?
273            .map_err(|e| Error::new(ErrorCode::CdpError, e.to_string()))?;
274        Ok(value)
275    }
276
277    /// Wait for a CDP event matching `predicate` (true = matches).
278    /// Returns `cdp_timeout` if `timeout` elapses first.
279    pub async fn wait_event<F>(
280        &self,
281        timeout: Duration,
282        mut predicate: F,
283    ) -> Result<CdpEvent, Error>
284    where
285        F: FnMut(&CdpEvent) -> bool,
286    {
287        let mut rx = self.events_tx.subscribe();
288        let deadline = tokio::time::Instant::now() + timeout;
289        loop {
290            let remaining = deadline.saturating_duration_since(tokio::time::Instant::now());
291            if remaining.is_zero() {
292                return Err(Error::new(ErrorCode::CdpTimeout, "wait_event: timed out"));
293            }
294            match tokio::time::timeout(remaining, rx.recv()).await {
295                Ok(Ok(ev)) if predicate(&ev) => return Ok(ev),
296                Ok(Ok(_)) => continue,
297                Ok(Err(broadcast::error::RecvError::Lagged(_))) => continue,
298                Ok(Err(broadcast::error::RecvError::Closed)) => {
299                    return Err(Error::new(
300                        ErrorCode::CdpUnavailable,
301                        "wait_event: events channel closed",
302                    ));
303                }
304                Err(_) => {
305                    return Err(Error::new(ErrorCode::CdpTimeout, "wait_event: timed out"));
306                }
307            }
308        }
309    }
310
311    pub fn close(&self) {
312        let _ = self.tx.send(OutMsg::Close);
313    }
314}
315
316impl Drop for Connection {
317    fn drop(&mut self) {
318        let _ = self.tx.send(OutMsg::Close);
319    }
320}
321
322async fn dispatch(msg: Value, pending: &PendingMap, events: &broadcast::Sender<CdpEvent>) {
323    if let Some(id) = msg.get("id").and_then(|v| v.as_i64()) {
324        let mut map = pending.lock().await;
325        if let Some(sender) = map.remove(&id) {
326            if let Some(err) = msg.get("error") {
327                let code = err.get("code").and_then(|v| v.as_i64()).unwrap_or(-1);
328                let message = err
329                    .get("message")
330                    .and_then(|v| v.as_str())
331                    .unwrap_or("")
332                    .to_string();
333                let _ = sender.send(Err(CdpRemoteError { code, message }));
334            } else {
335                let result = msg.get("result").cloned().unwrap_or(Value::Null);
336                let _ = sender.send(Ok(result));
337            }
338        }
339    } else if let Some(method) = msg.get("method").and_then(|v| v.as_str()) {
340        let params = msg.get("params").cloned().unwrap_or(Value::Null);
341        let session_id = msg
342            .get("sessionId")
343            .and_then(|v| v.as_str())
344            .map(str::to_string);
345        let _ = events.send(CdpEvent {
346            method: method.to_string(),
347            session_id,
348            params,
349        });
350    }
351}
352
353fn append_query_pairs(url: &str, pairs: &[(&str, &str)]) -> Result<String, Error> {
354    let mut parsed = url::Url::parse(url)
355        .map_err(|e| Error::new(ErrorCode::InvalidEndpoint, format!("CDP url {url:?}: {e}")))?;
356    {
357        let mut query = parsed.query_pairs_mut();
358        for (key, value) in pairs {
359            query.append_pair(key, value);
360        }
361    }
362    Ok(parsed.to_string())
363}
364
365fn build_ws_request(url: &str, _token: Option<&str>) -> Result<Request<()>, Error> {
366    // tokio_tungstenite::connect_async accepts &str directly via IntoClientRequest,
367    // but we go through the explicit Request type so we can attach headers later.
368    let uri: Uri = url
369        .parse()
370        .map_err(|e| Error::new(ErrorCode::InvalidEndpoint, format!("CDP url {url:?}: {e}")))?;
371    let host = uri.authority().map(|a| a.as_str()).unwrap_or("localhost");
372    let req = Request::builder()
373        .method("GET")
374        .uri(url)
375        .header(header::HOST, host)
376        .header(header::CONNECTION, "Upgrade")
377        .header(header::UPGRADE, "websocket")
378        .header(header::SEC_WEBSOCKET_VERSION, "13")
379        .header(header::SEC_WEBSOCKET_KEY, generate_key())
380        .body(())
381        .map_err(|e| Error::new(ErrorCode::InternalError, format!("CDP build request: {e}")))?;
382    // _token is already baked into the URL by the caller; bearer header
383    // would also be acceptable but axum's WebSocketUpgrade ignores it for
384    // the upgrade handshake.
385    req.into_client_request().map_err(|e| {
386        Error::new(
387            ErrorCode::InternalError,
388            format!("CDP into_client_request: {e}"),
389        )
390    })
391}
392
393#[cfg(test)]
394mod tests {
395    use super::*;
396
397    #[test]
398    fn query_pairs_are_percent_encoded() {
399        let url = append_query_pairs("ws://localhost:9222/cdp", &[("token", "a+b&c%20")]).unwrap();
400        assert_eq!(url, "ws://localhost:9222/cdp?token=a%2Bb%26c%2520");
401    }
402
403    // The bearer travels as `?token_secret=`; AFDATA redaction must scrub it
404    // before the failed-connect URL reaches the error envelope.
405    #[tokio::test]
406    async fn connect_error_redacts_token_secret() {
407        // Port 1 refuses fast, so we exercise the map_err path deterministically.
408        let err = Connection::connect("ws://127.0.0.1:1/cdp", Some("supersecret"))
409            .await
410            .err()
411            .expect("connect to closed port must fail");
412        let msg = err.to_string();
413        assert!(
414            msg.contains("token_secret=***"),
415            "token not redacted: {msg}"
416        );
417        assert!(!msg.contains("supersecret"), "raw token leaked: {msg}");
418    }
419}