Skip to main content

rightkit_control/
webdriver.rs

1//! W3C-WebDriver-shaped session bridge over a native [`Backend`].
2//!
3//! Wire-compatible subset: status, new/delete session, find element(s), element
4//! click/value/text, screenshot. On macOS the backend resolves elements through
5//! Accessibility and acts with background per-pid input. A proxy backend that
6//! forwards to an in-app embedded WebDriver (tauri-plugin-wdio) and only
7//! overrides input is specified in the design doc.
8use crate::admission::{CancelToken, EffectGate, EffectKind, EffectRequest, ExecError};
9use crate::json::{esc, str_field};
10use std::collections::HashMap;
11use std::io::{Read, Write};
12use std::net::TcpListener;
13use std::sync::Arc;
14use std::time::Duration;
15
16pub const ELEMENT_KEY: &str = "element-6066-11e4-a52e-4f735466cecf";
17
18pub trait Backend {
19    fn find(&mut self, using: &str, value: &str) -> Result<String, String>;
20    fn click(&mut self, id: &str) -> Result<(), String>;
21    fn send_keys(&mut self, id: &str, text: &str) -> Result<(), String>;
22    fn text(&mut self, id: &str) -> Result<String, String>;
23    fn screenshot_png_b64(&mut self) -> Result<String, String>;
24}
25
26pub struct Server<B: Backend> {
27    pub backend: B,
28    sessions: HashMap<String, ()>,
29    next: u32,
30    gate: Arc<EffectGate>,
31}
32
33pub fn b64(data: &[u8]) -> String {
34    const T: &[u8] = b"ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz0123456789+/";
35    let mut o = String::new();
36    for c in data.chunks(3) {
37        let n = (c[0] as u32) << 16
38            | (*c.get(1).unwrap_or(&0) as u32) << 8
39            | *c.get(2).unwrap_or(&0) as u32;
40        for i in 0..4 {
41            if i <= c.len() {
42                o.push(T[(n >> (18 - 6 * i) & 63) as usize] as char);
43            } else {
44                o.push('=');
45            }
46        }
47    }
48    o
49}
50
51fn ok(v: &str) -> (u16, String) {
52    (200, format!("{{\"value\":{v}}}"))
53}
54fn err(status: u16, code: &str, msg: &str) -> (u16, String) {
55    (
56        status,
57        format!(
58            "{{\"value\":{{\"error\":\"{code}\",\"message\":\"{}\",\"stacktrace\":\"\"}}}}",
59            esc(msg)
60        ),
61    )
62}
63
64impl<B: Backend> Server<B> {
65    pub fn new(backend: B) -> Self {
66        Self {
67            backend,
68            sessions: HashMap::new(),
69            next: 1,
70            gate: Arc::new(EffectGate::allow_all()),
71        }
72    }
73
74    /// Route every effectful command (element click / send keys / screenshot) through
75    /// `gate` (EFF-001). Default is [`EffectGate::allow_all`].
76    pub fn with_gate(mut self, gate: Arc<EffectGate>) -> Self {
77        self.gate = gate;
78        self
79    }
80
81    pub fn gate(&self) -> &Arc<EffectGate> {
82        &self.gate
83    }
84
85    fn effect<T>(
86        &mut self,
87        kind: EffectKind,
88        method: &str,
89        element: &str,
90        params: &str,
91        run: impl FnOnce(&mut B) -> Result<T, String>,
92    ) -> Result<T, String> {
93        let gate = self.gate.clone();
94        let backend = &mut self.backend;
95        let req = EffectRequest::new(kind, method, element).with_params(params);
96        let effect = gate.admit(req, &CancelToken::new(), |_| {
97            run(backend).map_err(ExecError::from)
98        });
99        match effect.settlement.outcome {
100            // Executor errors keep their pre-gate text; only gate refusals are prefixed.
101            crate::admission::SettleOutcome::Failed if !effect.settlement.crashed => {
102                Err(effect.settlement.reason.unwrap_or_default())
103            }
104            _ => effect.into_result(),
105        }
106    }
107
108    pub fn handle(&mut self, method: &str, path: &str, body: &str) -> (u16, String) {
109        let parts: Vec<&str> = path.trim_matches('/').split('/').collect();
110        match (method, parts.as_slice()) {
111            ("GET", ["status"]) => ok("{\"ready\":true,\"message\":\"rightkit-control\"}"),
112            ("POST", ["session"]) => {
113                let id = format!("rkc-{}", self.next);
114                self.next += 1;
115                self.sessions.insert(id.clone(), ());
116                ok(&format!(
117                    "{{\"sessionId\":\"{id}\",\"capabilities\":{{\"browserName\":\"tauri\",\"rightkit:input\":\"background-native\"}}}}"
118                ))
119            }
120            ("DELETE", ["session", id]) => {
121                self.sessions.remove(*id);
122                ok("null")
123            }
124            (_, ["session", id, rest @ ..]) if !self.sessions.contains_key(*id) => {
125                let _ = rest;
126                err(404, "invalid session id", "unknown session")
127            }
128            ("POST", ["session", _, "element"]) => {
129                let using = str_field(body, "using").unwrap_or_default();
130                let value = str_field(body, "value").unwrap_or_default();
131                match self.backend.find(&using, &value) {
132                    Ok(e) => ok(&format!("{{\"{ELEMENT_KEY}\":\"{}\"}}", esc(&e))),
133                    Err(m) => err(404, "no such element", &m),
134                }
135            }
136            ("POST", ["session", _, "element", e, "click"]) => {
137                match self.effect(EffectKind::Input, "click", e, body, |b| b.click(e)) {
138                    Ok(()) => ok("null"),
139                    Err(m) => err(400, "element click intercepted", &m),
140                }
141            }
142            ("POST", ["session", _, "element", e, "value"]) => {
143                let text = str_field(body, "text").unwrap_or_default();
144                match self.effect(EffectKind::Input, "send_keys", e, body, |b| {
145                    b.send_keys(e, &text)
146                }) {
147                    Ok(()) => ok("null"),
148                    Err(m) => err(400, "element not interactable", &m),
149                }
150            }
151            ("GET", ["session", _, "element", e, "text"]) => match self.backend.text(e) {
152                Ok(t) => ok(&format!("\"{}\"", esc(&t))),
153                Err(m) => err(404, "stale element reference", &m),
154            },
155            ("GET", ["session", _, "screenshot"]) => {
156                match self.effect(EffectKind::Capture, "screenshot", "window", "", |b| {
157                    b.screenshot_png_b64()
158                }) {
159                    Ok(b) => ok(&format!("\"{b}\"")),
160                    Err(m) => err(500, "unknown error", &m),
161                }
162            }
163            _ => err(404, "unknown command", &format!("{method} {path}")),
164        }
165    }
166
167    /// Blocking loopback server for third-party WebDriver clients. Every request
168    /// must pass [`check_request`]: loopback peer, a `Host` that names this
169    /// listener, no browser `Origin`, and the per-launch bearer credential. There
170    /// is no unauthenticated mode: loopback alone does not identify the caller.
171    pub fn serve(&mut self, listener: TcpListener, credential: &str) -> std::io::Result<()> {
172        if credential.len() < 32 {
173            return Err(std::io::Error::new(
174                std::io::ErrorKind::InvalidInput,
175                "credential must be at least 32 characters",
176            ));
177        }
178        let port = listener.local_addr()?.port();
179        for conn in listener.incoming() {
180            let mut s = match conn {
181                Ok(s) => s,
182                Err(_) => continue,
183            };
184            let _ = s.set_read_timeout(Some(Duration::from_secs(5)));
185            let _ = s.set_write_timeout(Some(Duration::from_secs(5)));
186            let peer_ok = s.peer_addr().map(|a| a.ip().is_loopback()).unwrap_or(false);
187            let (code, out) = match read_request(&mut s) {
188                Some((head, body)) => match check_request(&head, peer_ok, port, credential) {
189                    Ok(()) => {
190                        let mut first = head.lines().next().unwrap_or("").split(' ');
191                        let (m, p) = (first.next().unwrap_or(""), first.next().unwrap_or("/"));
192                        self.handle(m, p, &body)
193                    }
194                    Err(e) => e,
195                },
196                None => err(400, "invalid argument", "malformed request"),
197            };
198            let _ = write!(s, "HTTP/1.1 {code} X\r\nContent-Type: application/json\r\nContent-Length: {}\r\nConnection: close\r\n\r\n{out}", out.len());
199        }
200        Ok(())
201    }
202}
203
204const MAX_REQUEST: usize = 1 << 20;
205
206/// Read one HTTP request (head, body), bounded; `None` on malformed or oversized input.
207fn read_request(s: &mut std::net::TcpStream) -> Option<(String, String)> {
208    let mut buf: Vec<u8> = Vec::new();
209    let mut chunk = [0u8; 8192];
210    let (head_end, len) = loop {
211        let n = s.read(&mut chunk).ok()?;
212        if n == 0 {
213            return None;
214        }
215        buf.extend_from_slice(&chunk[..n]);
216        if buf.len() > MAX_REQUEST {
217            return None;
218        }
219        if let Some(i) = buf.windows(4).position(|w| w == b"\r\n\r\n") {
220            let head = String::from_utf8_lossy(&buf[..i]).to_string();
221            let cl = header(&head, "content-length")
222                .and_then(|v| v.parse::<usize>().ok())
223                .unwrap_or(0);
224            if cl > MAX_REQUEST {
225                return None;
226            }
227            break (i + 4, cl);
228        }
229    };
230    while buf.len() < head_end + len {
231        let n = s.read(&mut chunk).ok()?;
232        if n == 0 {
233            return None;
234        }
235        buf.extend_from_slice(&chunk[..n]);
236        if buf.len() > MAX_REQUEST + 8192 {
237            return None;
238        }
239    }
240    let head = String::from_utf8_lossy(&buf[..head_end - 4]).to_string();
241    let body = String::from_utf8_lossy(&buf[head_end..head_end + len]).to_string();
242    Some((head, body))
243}
244
245fn header(head: &str, name: &str) -> Option<String> {
246    head.lines().skip(1).find_map(|l| {
247        let (k, v) = l.split_once(':')?;
248        k.trim()
249            .eq_ignore_ascii_case(name)
250            .then(|| v.trim().to_string())
251    })
252}
253
254fn ct_eq(a: &str, b: &str) -> bool {
255    let (a, b) = (a.as_bytes(), b.as_bytes());
256    let mut d = u32::from(a.len() != b.len());
257    for i in 0..a.len().max(b.len()) {
258        d |= u32::from(a.get(i).copied().unwrap_or(0) ^ b.get(i).copied().unwrap_or(0));
259    }
260    d == 0
261}
262
263/// Admission for one request head. Pure so it is testable without a socket.
264pub fn check_request(
265    head: &str,
266    peer_is_loopback: bool,
267    port: u16,
268    credential: &str,
269) -> Result<(), (u16, String)> {
270    if !peer_is_loopback {
271        return Err(err(403, "unknown error", "peer is not loopback"));
272    }
273    let host_ok = header(head, "host")
274        .map(|h| {
275            h == format!("127.0.0.1:{port}")
276                || h == format!("localhost:{port}")
277                || h == format!("[::1]:{port}")
278        })
279        .unwrap_or(false);
280    if !host_ok {
281        return Err(err(403, "unknown error", "unexpected Host header"));
282    }
283    if header(head, "origin").is_some() {
284        return Err(err(
285            403,
286            "unknown error",
287            "browser-originated requests are refused",
288        ));
289    }
290    let presented = header(head, "authorization")
291        .and_then(|v| v.strip_prefix("Bearer ").map(str::to_string))
292        .unwrap_or_default();
293    if !ct_eq(&presented, credential) {
294        return Err(err(
295            401,
296            "unknown error",
297            "missing or wrong bridge credential",
298        ));
299    }
300    Ok(())
301}
302
303#[cfg(target_os = "macos")]
304pub mod macos {
305    use super::Backend;
306    use crate::mac::{self, Node};
307    use std::collections::HashMap;
308
309    /// Accessibility-resolved elements, background real-input actions.
310    ///
311    /// This is an adapter: its [`Backend`] methods act directly. Drive it
312    /// through a gated [`super::Server`] (`Server::new(AxBackend::new(pid))
313    /// .with_gate(gate)`), which admits click / send keys / screenshot.
314    pub struct AxBackend {
315        pub pid: i32,
316        els: HashMap<String, Node>,
317        n: u32,
318    }
319    impl AxBackend {
320        pub fn new(pid: i32) -> Self {
321            Self {
322                pid,
323                els: HashMap::new(),
324                n: 0,
325            }
326        }
327        fn get(&self, id: &str) -> Result<&Node, String> {
328            self.els
329                .get(id)
330                .ok_or_else(|| "unknown element".to_string())
331        }
332    }
333    impl Backend for AxBackend {
334        fn find(&mut self, using: &str, value: &str) -> Result<String, String> {
335            let exact = using == "link text" || using == "name";
336            let node = mac::ax_find(self.pid, value, exact)
337                .ok_or_else(|| format!("no element matching {value:?}"))?;
338            self.n += 1;
339            let id = format!("ax-{}", self.n);
340            self.els.insert(id.clone(), node);
341            Ok(id)
342        }
343        fn click(&mut self, id: &str) -> Result<(), String> {
344            let n = self.get(id)?.clone();
345            // AXPress first: works on inactive apps whose WKWebView swallows
346            // first-mouse clicks. Fall back to a real per-pid mouse click.
347            if mac::ax_press(&n) {
348                return Ok(());
349            }
350            mac::click_node(self.pid, &n)
351        }
352        fn send_keys(&mut self, id: &str, text: &str) -> Result<(), String> {
353            let n = self.get(id)?.clone();
354            mac::click_node(self.pid, &n)?; // focus inside the app, not the OS
355            mac::type_text(self.pid, text);
356            Ok(())
357        }
358        fn text(&mut self, id: &str) -> Result<String, String> {
359            Ok(mac::ax_text(self.get(id)?))
360        }
361        fn screenshot_png_b64(&mut self) -> Result<String, String> {
362            let w = mac::main_window(self.pid).ok_or("no window")?;
363            let p = std::env::temp_dir().join(format!("rkc-{}.png", std::process::id()));
364            if !mac::capture_window(w.id, &p.to_string_lossy()) {
365                return Err("screencapture failed".into());
366            }
367            let bytes = std::fs::read(&p).map_err(|e| e.to_string())?;
368            let _ = std::fs::remove_file(&p);
369            Ok(super::b64(&bytes))
370        }
371    }
372}
373
374#[cfg(test)]
375mod tests {
376    use super::*;
377    struct Fake(Vec<String>);
378    impl Backend for Fake {
379        fn find(&mut self, u: &str, v: &str) -> Result<String, String> {
380            self.0.push(format!("find {u} {v}"));
381            Ok("e1".into())
382        }
383        fn click(&mut self, id: &str) -> Result<(), String> {
384            self.0.push(format!("click {id}"));
385            Ok(())
386        }
387        fn send_keys(&mut self, _: &str, t: &str) -> Result<(), String> {
388            self.0.push(format!("keys {t}"));
389            Ok(())
390        }
391        fn text(&mut self, _: &str) -> Result<String, String> {
392            Ok("hi".into())
393        }
394        fn screenshot_png_b64(&mut self) -> Result<String, String> {
395            Ok(b64(b"abcd"))
396        }
397    }
398    #[test]
399    fn session_flow() {
400        let mut s = Server::new(Fake(vec![]));
401        let (_, r) = s.handle("POST", "/session", "{}");
402        assert!(r.contains("rkc-1"));
403        let (c, r) = s.handle(
404            "POST",
405            "/session/rkc-1/element",
406            r#"{"using":"link text","value":"Go"}"#,
407        );
408        assert_eq!(c, 200);
409        assert!(r.contains(ELEMENT_KEY));
410        assert_eq!(
411            s.handle("POST", "/session/rkc-1/element/e1/click", "{}").0,
412            200
413        );
414        assert_eq!(s.handle("POST", "/session/zzz/element", "{}").0, 404);
415        assert_eq!(s.backend.0, ["find link text Go", "click e1"]);
416        assert_eq!(b64(b"abcd"), "YWJjZA==");
417    }
418
419    #[test]
420    fn denied_click_never_reaches_backend() {
421        use crate::admission::{AdmissionDecision, SettleOutcome};
422        let gate = Arc::new(EffectGate::new(|r: &EffectRequest| {
423            if r.method == "click" {
424                AdmissionDecision::deny("clicks blocked")
425            } else {
426                AdmissionDecision::Allow
427            }
428        }));
429        let mut s = Server::new(Fake(vec![])).with_gate(gate.clone());
430        s.handle("POST", "/session", "{}");
431        s.handle(
432            "POST",
433            "/session/rkc-1/element",
434            r#"{"using":"css","value":"x"}"#,
435        );
436        let (c, r) = s.handle("POST", "/session/rkc-1/element/e1/click", "{}");
437        assert_eq!(c, 400);
438        assert!(r.contains("denied"), "{r}");
439        let (c, _) = s.handle(
440            "POST",
441            "/session/rkc-1/element/e1/value",
442            r#"{"text":"hi"}"#,
443        );
444        assert_eq!(c, 200);
445        assert_eq!(s.backend.0, ["find css x", "keys hi"]);
446        let outcomes: Vec<_> = gate.settlements().iter().map(|x| x.outcome).collect();
447        assert_eq!(outcomes, [SettleOutcome::Denied, SettleOutcome::Ok]);
448    }
449}