Skip to main content

browser_commander/browser/webdriver/
bidi.rs

1//! Native WebDriver BiDi multiplexing with typed context and script operations.
2
3use anyhow::{anyhow, Result};
4use async_tungstenite::tungstenite::{protocol::WebSocketConfig, Message};
5use futures::StreamExt;
6use serde::{Deserialize, Serialize};
7use serde_json::{json, Value};
8use std::{
9    collections::HashMap,
10    sync::{
11        atomic::{AtomicBool, AtomicU64, Ordering},
12        Arc, Mutex,
13    },
14    time::Duration,
15};
16use tokio::{
17    sync::{broadcast, mpsc, oneshot},
18    task::JoinHandle,
19    time::timeout,
20};
21
22#[derive(Clone, Debug, Deserialize, Serialize)]
23pub struct BidiEvent {
24    pub method: String,
25    pub params: Value,
26}
27
28#[derive(Clone, Debug, Deserialize, Serialize)]
29#[serde(rename_all = "camelCase")]
30pub struct BrowsingContext {
31    pub context: String,
32    pub url: String,
33    pub children: Option<Vec<BrowsingContext>>,
34    #[serde(default)]
35    pub parent: Option<String>,
36}
37
38#[derive(Clone, Debug, Deserialize, Serialize)]
39pub struct ContextTree {
40    pub contexts: Vec<BrowsingContext>,
41}
42
43#[derive(Clone, Debug, Deserialize, Serialize)]
44pub struct NavigationResult {
45    pub navigation: Option<String>,
46    pub url: String,
47}
48
49#[derive(Clone, Debug, Deserialize, Serialize)]
50#[serde(tag = "type", rename_all = "lowercase")]
51pub enum ScriptResult {
52    Success {
53        result: Value,
54        realm: String,
55    },
56    Exception {
57        #[serde(rename = "exceptionDetails")]
58        exception_details: Value,
59        realm: String,
60    },
61}
62
63struct State {
64    closed: AtomicBool,
65    pending: Mutex<HashMap<u64, oneshot::Sender<Result<Value>>>>,
66    events: broadcast::Sender<BidiEvent>,
67}
68impl State {
69    fn disconnect(&self) {
70        if self.closed.swap(true, Ordering::AcqRel) {
71            return;
72        }
73        for (_, sender) in self.pending.lock().unwrap().drain() {
74            let _ = sender.send(Err(anyhow!("BiDi disconnected")));
75        }
76        let _ = self.events.send(BidiEvent {
77            method: "disconnected".into(),
78            params: json!({}),
79        });
80    }
81}
82struct Pending {
83    id: u64,
84    state: Arc<State>,
85}
86impl Drop for Pending {
87    fn drop(&mut self) {
88        self.state.pending.lock().unwrap().remove(&self.id);
89    }
90}
91
92/// One owned connection. Up to 64 outstanding calls, 4 MiB incoming messages,
93/// and 1,024 retained events. Slow subscribers receive `RecvError::Lagged`.
94pub struct BidiClient {
95    state: Arc<State>,
96    next: AtomicU64,
97    outgoing: mpsc::Sender<Value>,
98    shutdown: tokio::sync::watch::Sender<bool>,
99    task: Mutex<Option<JoinHandle<()>>>,
100}
101impl BidiClient {
102    pub async fn connect(url: &str) -> Result<Self> {
103        let parsed = url::Url::parse(url)?;
104        if !matches!(parsed.scheme(), "ws" | "wss")
105            || !matches!(
106                parsed.host_str(),
107                Some("127.0.0.1" | "localhost" | "[::1]" | "::1")
108            )
109        {
110            return Err(anyhow!(
111                "managed BiDi endpoint must be a loopback WebSocket"
112            ));
113        }
114        // The native driver is local. Bound handshake latency as well as frames.
115        let config = WebSocketConfig::default()
116            .max_message_size(Some(4 * 1024 * 1024))
117            .max_frame_size(Some(4 * 1024 * 1024));
118        let (mut socket, _) = timeout(
119            Duration::from_secs(10),
120            async_tungstenite::tokio::connect_async_with_config(url, Some(config)),
121        )
122        .await??;
123        let (outgoing, mut queue) = mpsc::channel::<Value>(64);
124        let (shutdown, mut closing) = tokio::sync::watch::channel(false);
125        let state = Arc::new(State {
126            closed: AtomicBool::new(false),
127            pending: Mutex::new(HashMap::new()),
128            events: broadcast::channel(1024).0,
129        });
130        let actor = state.clone();
131        let task = tokio::spawn(async move {
132            loop {
133                tokio::select! {
134                    _=closing.changed()=>break,
135                    request=queue.recv()=>match request {
136                        Some(request)=>if socket.send(Message::Text(request.to_string().into())).await.is_err(){break;},
137                        None=>break,
138                    },
139                    frame=socket.next()=>match frame {
140                        Some(Ok(Message::Text(text)))=>{
141                            let Ok(message)=serde_json::from_str::<Value>(&text) else {tracing::debug!("Ignoring malformed BiDi message");continue;};
142                            if let Some(id)=message.get("id").and_then(Value::as_u64) {
143                                if let Some(sender)=actor.pending.lock().unwrap().remove(&id) {
144                                    let result=if message.get("type").and_then(Value::as_str)==Some("error") || message.get("error").is_some() {
145                                        Err(anyhow!("BiDi {}: {}",message.get("error").and_then(Value::as_str).unwrap_or("error"),message.get("message").and_then(Value::as_str).unwrap_or("unknown error")))
146                                    } else {Ok(message.get("result").cloned().unwrap_or(Value::Null))};
147                                    let _=sender.send(result);
148                                }
149                            } else if let Some(method)=message.get("method").and_then(Value::as_str) {
150                                let _=actor.events.send(BidiEvent{method:method.into(),params:message.get("params").cloned().unwrap_or_else(||json!({}))});
151                            }
152                        },
153                        Some(Ok(Message::Ping(data)))=>if socket.send(Message::Pong(data)).await.is_err(){break;},
154                        None|Some(Err(_))|Some(Ok(Message::Close(_)))=>break,
155                        _=>{},
156                    }
157                }
158            }
159            actor.disconnect();
160            let _ = timeout(Duration::from_secs(1), socket.close(None)).await;
161        });
162        Ok(Self {
163            state,
164            next: AtomicU64::new(1),
165            outgoing,
166            shutdown,
167            task: Mutex::new(Some(task)),
168        })
169    }
170
171    /// Dynamic fallback for newly introduced BiDi protocol commands.
172    pub async fn send(&self, method: &str, params: Value) -> Result<Value> {
173        let id = self.next.fetch_add(1, Ordering::Relaxed);
174        let (sender, receiver) = oneshot::channel();
175        {
176            let mut pending = self.state.pending.lock().unwrap();
177            if self.state.closed.load(Ordering::Acquire) {
178                return Err(anyhow!("BiDi disconnected"));
179            }
180            if pending.len() >= 64 {
181                return Err(anyhow!("Too many pending BiDi requests"));
182            }
183            pending.insert(id, sender);
184        }
185        let _guard = Pending {
186            id,
187            state: self.state.clone(),
188        };
189        timeout(Duration::from_secs(30), async {
190            self.outgoing
191                .send(json!({"id":id,"method":method,"params":params}))
192                .await
193                .map_err(|_| anyhow!("BiDi disconnected"))?;
194            receiver.await.map_err(|_| anyhow!("BiDi disconnected"))?
195        })
196        .await
197        .map_err(|_| anyhow!("BiDi request timed out"))?
198    }
199    pub fn events(&self) -> broadcast::Receiver<BidiEvent> {
200        self.state.events.subscribe()
201    }
202    pub async fn subscribe(&self, events: &[&str], contexts: Option<&[&str]>) -> Result<()> {
203        let mut params = json!({"events":events});
204        if let Some(contexts) = contexts {
205            params["contexts"] = json!(contexts);
206        }
207        self.send("session.subscribe", params).await?;
208        Ok(())
209    }
210    pub async fn unsubscribe(&self, events: &[&str], contexts: Option<&[&str]>) -> Result<()> {
211        let mut params = json!({"events":events});
212        if let Some(contexts) = contexts {
213            params["contexts"] = json!(contexts);
214        }
215        self.send("session.unsubscribe", params).await?;
216        Ok(())
217    }
218    pub async fn get_tree(&self) -> Result<ContextTree> {
219        Ok(serde_json::from_value(
220            self.send("browsingContext.getTree", json!({})).await?,
221        )?)
222    }
223    pub async fn navigate(&self, context: &str, url: &str) -> Result<NavigationResult> {
224        Ok(serde_json::from_value(
225            self.send(
226                "browsingContext.navigate",
227                json!({"context":context,"url":url,"wait":"complete"}),
228            )
229            .await?,
230        )?)
231    }
232    pub async fn evaluate(&self, context: &str, expression: &str) -> Result<ScriptResult> {
233        Ok(serde_json::from_value(self.send("script.evaluate",json!({"target":{"context":context},"expression":expression,"awaitPromise":true,"resultOwnership":"none"})).await?)?)
234    }
235    pub async fn close(&self) {
236        let task = self.task.lock().unwrap().take();
237        self.state.disconnect();
238        let _ = self.shutdown.send(true);
239        if let Some(mut task) = task {
240            if timeout(Duration::from_secs(2), &mut task).await.is_err() {
241                task.abort();
242                let _ = task.await;
243            }
244        }
245    }
246}
247impl Drop for BidiClient {
248    fn drop(&mut self) {
249        if let Some(task) = self.task.lock().unwrap().take() {
250            task.abort();
251        }
252        self.state.disconnect();
253    }
254}