1use 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
33pub 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 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 pub async fn connect(endpoint_ws_url: &str, token: Option<&str>) -> Result<Self, Error> {
81 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)?;
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 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 let (ws, _resp) = tokio_tungstenite::connect_async(request)
133 .await
134 .map_err(|e| {
135 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)?;
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 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 pub fn subscribe(&self) -> broadcast::Receiver<CdpEvent> {
234 self.events_tx.subscribe()
235 }
236
237 pub async fn send<P: Serialize>(
240 &self,
241 method: &str,
242 params: &P,
243 session_id: Option<&str>,
244 ) -> Result<Value, Error> {
245 self.send_inner(method, params, session_id, None).await
246 }
247
248 pub async fn send_timeout<P: Serialize>(
252 &self,
253 method: &str,
254 params: &P,
255 session_id: Option<&str>,
256 timeout: Duration,
257 ) -> Result<Value, Error> {
258 self.send_inner(method, params, session_id, Some(timeout))
259 .await
260 }
261
262 async fn send_inner<P: Serialize>(
263 &self,
264 method: &str,
265 params: &P,
266 session_id: Option<&str>,
267 timeout: Option<Duration>,
268 ) -> Result<Value, Error> {
269 let id = self.next_id.fetch_add(1, Ordering::SeqCst);
270 let body = match session_id {
271 Some(sid) => serde_json::json!({
272 "id": id,
273 "method": method,
274 "params": params,
275 "sessionId": sid,
276 }),
277 None => serde_json::json!({
278 "id": id,
279 "method": method,
280 "params": params,
281 }),
282 };
283 let serialized = serde_json::to_string(&body).map_err(|e| {
284 Error::new(
285 ErrorCode::InternalError,
286 format!("CDP send: serialize {method}: {e}"),
287 )
288 })?;
289 let (resp_tx, resp_rx) = oneshot::channel();
290 self.pending.lock().await.insert(id, resp_tx);
291 self.tx
292 .send(OutMsg::Text(serialized))
293 .map_err(|_| Error::new(ErrorCode::CdpUnavailable, "CDP writer closed before send"))?;
294 let received = if let Some(timeout) = timeout {
295 match tokio::time::timeout(timeout, resp_rx).await {
296 Ok(value) => value,
297 Err(_) => {
298 self.pending.lock().await.remove(&id);
299 return Err(Error::new(
300 ErrorCode::CdpTimeout,
301 format!("{method}: CDP reply timed out after {timeout:?}"),
302 ));
303 }
304 }
305 } else {
306 resp_rx.await
307 };
308 let value = received
309 .map_err(|_| Error::new(ErrorCode::CdpUnavailable, "CDP reader closed"))?
310 .map_err(|e| Error::new(ErrorCode::CdpError, e.to_string()))?;
311 Ok(value)
312 }
313
314 pub async fn wait_event<F>(
317 &self,
318 timeout: Duration,
319 mut predicate: F,
320 ) -> Result<CdpEvent, Error>
321 where
322 F: FnMut(&CdpEvent) -> bool,
323 {
324 let mut rx = self.events_tx.subscribe();
325 let deadline = tokio::time::Instant::now() + timeout;
326 loop {
327 let remaining = deadline.saturating_duration_since(tokio::time::Instant::now());
328 if remaining.is_zero() {
329 return Err(Error::new(ErrorCode::CdpTimeout, "wait_event: timed out"));
330 }
331 match tokio::time::timeout(remaining, rx.recv()).await {
332 Ok(Ok(ev)) if predicate(&ev) => return Ok(ev),
333 Ok(Ok(_)) => continue,
334 Ok(Err(broadcast::error::RecvError::Lagged(_))) => continue,
335 Ok(Err(broadcast::error::RecvError::Closed)) => {
336 return Err(Error::new(
337 ErrorCode::CdpUnavailable,
338 "wait_event: events channel closed",
339 ));
340 }
341 Err(_) => {
342 return Err(Error::new(ErrorCode::CdpTimeout, "wait_event: timed out"));
343 }
344 }
345 }
346 }
347
348 pub fn close(&self) {
349 let _ = self.tx.send(OutMsg::Close);
350 }
351}
352
353impl Drop for Connection {
354 fn drop(&mut self) {
355 let _ = self.tx.send(OutMsg::Close);
356 }
357}
358
359async fn dispatch(msg: Value, pending: &PendingMap, events: &broadcast::Sender<CdpEvent>) {
360 if let Some(id) = msg.get("id").and_then(|v| v.as_i64()) {
361 let mut map = pending.lock().await;
362 if let Some(sender) = map.remove(&id) {
363 if let Some(err) = msg.get("error") {
364 let code = err.get("code").and_then(|v| v.as_i64()).unwrap_or(-1);
365 let message = err
366 .get("message")
367 .and_then(|v| v.as_str())
368 .unwrap_or("")
369 .to_string();
370 let _ = sender.send(Err(CdpRemoteError { code, message }));
371 } else {
372 let result = msg.get("result").cloned().unwrap_or(Value::Null);
373 let _ = sender.send(Ok(result));
374 }
375 }
376 } else if let Some(method) = msg.get("method").and_then(|v| v.as_str()) {
377 let params = msg.get("params").cloned().unwrap_or(Value::Null);
378 let session_id = msg
379 .get("sessionId")
380 .and_then(|v| v.as_str())
381 .map(str::to_string);
382 let _ = events.send(CdpEvent {
383 method: method.to_string(),
384 session_id,
385 params,
386 });
387 }
388}
389
390fn append_query_pairs(url: &str, pairs: &[(&str, &str)]) -> Result<String, Error> {
391 let mut parsed = url::Url::parse(url)
392 .map_err(|e| Error::new(ErrorCode::InvalidEndpoint, format!("CDP url {url:?}: {e}")))?;
393 {
394 let mut query = parsed.query_pairs_mut();
395 for (key, value) in pairs {
396 query.append_pair(key, value);
397 }
398 }
399 Ok(parsed.to_string())
400}
401
402fn build_ws_request(url: &str) -> Result<Request<()>, Error> {
403 let uri: Uri = url
409 .parse()
410 .map_err(|e| Error::new(ErrorCode::InvalidEndpoint, format!("CDP url {url:?}: {e}")))?;
411 let host = uri.authority().map(|a| a.as_str()).unwrap_or("localhost");
412 let req = Request::builder()
413 .method("GET")
414 .uri(url)
415 .header(header::HOST, host)
416 .header(header::CONNECTION, "Upgrade")
417 .header(header::UPGRADE, "websocket")
418 .header(header::SEC_WEBSOCKET_VERSION, "13")
419 .header(header::SEC_WEBSOCKET_KEY, generate_key())
420 .body(())
421 .map_err(|e| Error::new(ErrorCode::InternalError, format!("CDP build request: {e}")))?;
422 req.into_client_request().map_err(|e| {
423 Error::new(
424 ErrorCode::InternalError,
425 format!("CDP into_client_request: {e}"),
426 )
427 })
428}
429
430#[cfg(test)]
431mod tests {
432 use super::*;
433
434 #[test]
435 fn query_pairs_are_percent_encoded() {
436 let url = append_query_pairs("ws://localhost:9222/cdp", &[("token", "a+b&c%20")]).unwrap();
437 assert_eq!(url, "ws://localhost:9222/cdp?token=a%2Bb%26c%2520");
438 }
439
440 #[tokio::test]
443 async fn connect_error_redacts_token_secret() {
444 let err = Connection::connect("ws://127.0.0.1:1/cdp", Some("supersecret"))
446 .await
447 .err()
448 .expect("connect to closed port must fail");
449 let msg = err.to_string();
450 assert!(
451 msg.contains("token_secret=***"),
452 "token not redacted: {msg}"
453 );
454 assert!(!msg.contains("supersecret"), "raw token leaked: {msg}");
455 }
456}