agent_first_http/sdk/cdp/
ws_client.rs1use 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, 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 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, 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 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 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 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 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 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 #[tokio::test]
406 async fn connect_error_redacts_token_secret() {
407 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}