browser_commander/browser/extension_relay/
mod.rs1mod protocol;
10mod server;
11mod state;
12
13pub use protocol::{extension_id_from_origin, DEFAULT_RELAY_PORT, RELAY_PATH};
14use serde::{Deserialize, Serialize};
15use serde_json::{json, Value};
16use state::{PendingGuard, SessionState, State};
17use std::{
18 sync::{atomic::Ordering, Arc},
19 time::Duration,
20};
21use tokio::{
22 net::TcpListener,
23 sync::{broadcast, oneshot},
24 task::JoinHandle,
25 time::timeout,
26};
27
28#[derive(Clone, Debug)]
29pub struct RelayOptions {
30 pub host: String,
31 pub port: u16,
32 pub timeout: Duration,
33 pub request_timeout: Duration,
34 pub allowed_extension_ids: Option<Vec<String>>,
35}
36
37impl Default for RelayOptions {
38 fn default() -> Self {
39 Self {
40 host: "127.0.0.1".into(),
41 port: DEFAULT_RELAY_PORT,
42 timeout: Duration::from_secs(60),
43 request_timeout: Duration::from_secs(30),
44 allowed_extension_ids: None,
45 }
46 }
47}
48
49#[derive(Debug, thiserror::Error)]
50pub enum RelayError {
51 #[error("{0}")]
52 InvalidOptions(String),
53 #[error("Extension relay I/O failed: {0}")]
54 Io(#[from] std::io::Error),
55 #[error("Extension relay response is invalid: {0}")]
56 Protocol(#[from] serde_json::Error),
57 #[error("{0}")]
58 Disconnected(String),
59 #[error("Extension request failed: {0}")]
60 Remote(String),
61 #[error("Extension relay timed out")]
62 Timeout,
63 #[error("Too many pending extension requests")]
64 Busy,
65}
66
67#[derive(Clone, Debug, Serialize, Deserialize)]
68#[serde(rename_all = "camelCase")]
69pub struct RelayExtension {
70 pub id: String,
71 pub version: String,
72 pub user_agent: String,
73}
74
75#[derive(Clone, Debug, Serialize, Deserialize)]
76#[serde(rename_all = "camelCase")]
77pub struct RelayTab {
78 pub tab_id: i64,
79 pub url: String,
80 pub title: String,
81 pub active: bool,
82}
83
84#[derive(Clone, Debug, Serialize, Deserialize)]
85pub struct RelayEvent {
86 pub method: String,
87 pub params: Value,
88}
89
90pub struct ExtensionRelay {
93 state: Arc<State>,
94 url: String,
95 port: u16,
96 task: Option<JoinHandle<()>>,
97}
98
99impl ExtensionRelay {
100 pub async fn listen(options: RelayOptions) -> Result<Self, RelayError> {
101 protocol::validate(&options)?;
102 let host = if options.host == "localhost" {
103 "127.0.0.1"
104 } else {
105 &options.host
106 };
107 let listener = TcpListener::bind((host, options.port)).await?;
108 let address = listener.local_addr()?;
109 let state = Arc::new(State::new(options));
110 let task = tokio::spawn(server::serve(listener, state.clone()));
111 Ok(Self {
112 state,
113 url: format!("ws://{address}{RELAY_PATH}"),
114 port: address.port(),
115 task: Some(task),
116 })
117 }
118
119 pub fn url(&self) -> &str {
120 &self.url
121 }
122 pub fn port(&self) -> u16 {
123 self.port
124 }
125 pub fn connected(&self) -> bool {
126 self.state.extension.borrow().is_some()
127 }
128 pub fn extension(&self) -> Option<RelayExtension> {
129 self.state.extension.borrow().clone()
130 }
131
132 pub async fn wait_for_extension(&self) -> Result<RelayExtension, RelayError> {
133 let mut extension = self.state.extension.subscribe();
134 let mut shutdown = self.state.shutdown.subscribe();
135 timeout(self.state.options.timeout, async {
136 loop {
137 if *shutdown.borrow() { return Err(RelayError::Disconnected("the relay was closed".into())); }
138 if let Some(info) = extension.borrow().clone() { return Ok(info); }
139 tokio::select! {
140 result = extension.changed() => { result.map_err(|_| RelayError::Disconnected("the relay was closed".into()))?; }
141 _ = shutdown.changed() => return Err(RelayError::Disconnected("the relay was closed".into())),
142 }
143 }
144 }).await.map_err(|_| RelayError::Timeout)?
145 }
146
147 pub async fn tabs(&self) -> Result<Vec<RelayTab>, RelayError> {
148 Ok(serde_json::from_value(
149 request(&self.state, "tabs.list", json!({})).await?,
150 )?)
151 }
152
153 pub async fn new_tab(&self, url: Option<&str>) -> Result<i64, RelayError> {
154 #[derive(Deserialize)]
155 #[serde(rename_all = "camelCase")]
156 struct Tab {
157 tab_id: i64,
158 }
159 let params = url
160 .map(|url| json!({"url":url}))
161 .unwrap_or_else(|| json!({}));
162 let tab: Tab = serde_json::from_value(request(&self.state, "tabs.create", params).await?)?;
163 Ok(tab.tab_id)
164 }
165
166 pub async fn session(&self, tab_id: i64) -> Result<RelaySession, RelayError> {
167 let _lock = self.state.session_lock.lock().await;
168 let existing = self
169 .state
170 .sessions
171 .lock()
172 .unwrap()
173 .get(&tab_id)
174 .and_then(std::sync::Weak::upgrade)
175 .filter(|session| !session.detached.load(Ordering::Acquire));
176 if let Some(state) = existing {
177 return Ok(RelaySession { state });
178 }
179 let session = Arc::new(SessionState {
180 tab_id,
181 detached: false.into(),
182 events: broadcast::channel(1024).0,
183 relay: self.state.clone(),
184 });
185 self.state
186 .sessions
187 .lock()
188 .unwrap()
189 .insert(tab_id, Arc::downgrade(&session));
190 struct AttachGuard(Arc<SessionState>, bool);
192 impl Drop for AttachGuard {
193 fn drop(&mut self) {
194 if !self.1 {
195 self.0.mark_detached("attach_failed");
196 }
197 }
198 }
199 let mut guard = AttachGuard(session.clone(), false);
200 request(&self.state, "debugger.attach", json!({"tabId":tab_id})).await?;
201 guard.1 = true;
202 Ok(RelaySession { state: session })
203 }
204
205 pub async fn close(&mut self) {
206 self.state.shutdown.send_replace(true);
207 if let Some(task) = self.task.take() {
208 let _ = task.await;
209 }
210 }
211}
212
213impl Drop for ExtensionRelay {
214 fn drop(&mut self) {
215 self.state.shutdown.send_replace(true);
216 self.state.disconnect("the relay was closed");
217 if let Some(task) = self.task.take() {
218 task.abort();
219 }
220 }
221}
222
223#[derive(Clone)]
224pub struct RelaySession {
225 state: Arc<SessionState>,
226}
227
228impl RelaySession {
229 pub fn tab_id(&self) -> i64 {
230 self.state.tab_id
231 }
232 pub fn is_detached(&self) -> bool {
233 self.state.detached.load(Ordering::Acquire)
234 }
235 pub fn subscribe(&self) -> broadcast::Receiver<RelayEvent> {
238 self.state.events.subscribe()
239 }
240 pub async fn send(&self, method: &str, params: Value) -> Result<Value, RelayError> {
241 if self.is_detached() {
242 return Err(RelayError::Disconnected(
243 "Debugger session is detached".into(),
244 ));
245 }
246 request(
247 &self.state.relay,
248 "cdp.send",
249 json!({"tabId":self.tab_id(),"method":method,"params":params}),
250 )
251 .await
252 }
253 pub async fn detach(&self) -> Result<(), RelayError> {
254 if !self.is_detached() {
255 request(
256 &self.state.relay,
257 "debugger.detach",
258 json!({"tabId":self.tab_id()}),
259 )
260 .await?;
261 self.state.mark_detached("detached_by_client");
262 }
263 Ok(())
264 }
265}
266
267#[async_trait::async_trait]
268impl crate::fingerprint::CdpTransport for RelaySession {
269 async fn send(&self, method: &str, params: Value) -> anyhow::Result<Value> {
270 Ok(RelaySession::send(self, method, params).await?)
271 }
272}
273
274async fn request(state: &Arc<State>, method: &str, params: Value) -> Result<Value, RelayError> {
275 let sender = state.outgoing.lock().unwrap().clone().ok_or_else(|| {
276 RelayError::Disconnected("The Browser Commander Relay extension is not connected".into())
277 })?;
278 let id = state.next_id.fetch_add(1, Ordering::Relaxed);
279 let (response, result) = oneshot::channel();
280 {
281 let mut pending = state.pending.lock().unwrap();
282 if pending.len() >= 64 {
283 return Err(RelayError::Busy);
284 }
285 pending.insert(id, response);
286 }
287 let _guard = PendingGuard {
288 id,
289 state: state.clone(),
290 };
291 sender
292 .try_send(json!({"id":id,"method":method,"params":params}))
293 .map_err(|error| match error {
294 tokio::sync::mpsc::error::TrySendError::Full(_) => RelayError::Busy,
295 _ => RelayError::Disconnected("the extension disconnected".into()),
296 })?;
297 timeout(state.options.request_timeout, result)
298 .await
299 .map_err(|_| RelayError::Timeout)?
300 .map_err(|_| RelayError::Disconnected("the extension disconnected".into()))?
301}
302
303pub async fn attach_via_extension(options: RelayOptions) -> Result<ExtensionRelay, RelayError> {
306 let relay = ExtensionRelay::listen(options).await?;
307 relay.wait_for_extension().await?;
308 Ok(relay)
309}
310
311pub fn write_extension_directory(
314 destination: impl AsRef<std::path::Path>,
315) -> Result<std::path::PathBuf, RelayError> {
316 let destination = destination.as_ref();
317 std::fs::create_dir_all(destination)?;
318 for (name, contents) in [
319 (
320 "manifest.json",
321 include_str!("extension_assets/manifest.json"),
322 ),
323 (
324 "background.js",
325 include_str!("extension_assets/background.js"),
326 ),
327 (
328 "relay-handler.js",
329 include_str!("extension_assets/relay-handler.js"),
330 ),
331 ] {
332 std::fs::write(destination.join(name), contents)?;
333 }
334 Ok(destination.to_owned())
335}
336
337#[cfg(test)]
338mod tests {
339 use super::*;
340
341 #[tokio::test]
342 async fn cancelled_requests_release_pending_slots_and_ignore_late_responses() {
343 let state = Arc::new(State::new(RelayOptions::default()));
344 let (sender, mut messages) = tokio::sync::mpsc::channel(64);
345 *state.outgoing.lock().unwrap() = Some(sender);
346 let task_state = state.clone();
347 let task = tokio::spawn(async move { request(&task_state, "tabs.list", json!({})).await });
348 let sent = timeout(Duration::from_secs(2), messages.recv())
349 .await
350 .unwrap()
351 .unwrap();
352 assert_eq!(state.pending.lock().unwrap().len(), 1);
353 task.abort();
354 assert!(task.await.unwrap_err().is_cancelled());
355 assert!(state.pending.lock().unwrap().is_empty());
356 state.message(json!({"id":sent["id"], "result":[]}));
357 assert!(state.pending.lock().unwrap().is_empty());
358 }
359
360 #[tokio::test]
361 async fn pending_request_count_is_bounded_and_disconnect_releases_all_slots() {
362 let state = Arc::new(State::new(RelayOptions::default()));
363 let (sender, mut messages) = tokio::sync::mpsc::channel(64);
364 *state.outgoing.lock().unwrap() = Some(sender);
365 let mut tasks = tokio::task::JoinSet::new();
366 for _ in 0..64 {
367 let state = state.clone();
368 tasks.spawn(async move { request(&state, "tabs.list", json!({})).await });
369 }
370 for _ in 0..64 {
371 timeout(Duration::from_secs(2), messages.recv())
372 .await
373 .unwrap()
374 .unwrap();
375 }
376 assert!(matches!(
377 request(&state, "tabs.list", json!({})).await,
378 Err(RelayError::Busy)
379 ));
380 state.disconnect("test disconnected");
381 while let Some(result) = tasks.join_next().await {
382 assert!(matches!(result.unwrap(), Err(RelayError::Disconnected(_))));
383 }
384 assert!(state.pending.lock().unwrap().is_empty());
385 }
386
387 #[test]
388 fn bundled_extension_matches_the_javascript_protocol() {
389 let directory =
390 crate::browser::profile_directory::create_temporary_user_data_dir(None).unwrap();
391 write_extension_directory(&directory).unwrap();
392 for name in ["manifest.json", "background.js", "relay-handler.js"] {
393 assert_eq!(
394 std::fs::read(directory.join(name)).unwrap(),
395 std::fs::read(
396 std::path::Path::new(env!("CARGO_MANIFEST_DIR"))
397 .join("../js/extension")
398 .join(name)
399 )
400 .unwrap()
401 );
402 }
403 std::fs::remove_dir_all(directory).unwrap();
404 }
405}