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) 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(
86        &mut self,
87        kind: EffectKind,
88        method: &str,
89        element: &str,
90        params: &str,
91        run: impl FnOnce(&mut B) -> Result<(), String>,
92    ) -> Result<(), 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"]) => match self.backend.screenshot_png_b64() {
156                Ok(b) => ok(&format!("\"{b}\"")),
157                Err(m) => err(500, "unknown error", &m),
158            },
159            _ => err(404, "unknown command", &format!("{method} {path}")),
160        }
161    }
162
163    /// Blocking loopback server for third-party WebDriver clients. Every request
164    /// must pass [`check_request`]: loopback peer, a `Host` that names this
165    /// listener, no browser `Origin`, and the per-launch bearer credential. There
166    /// is no unauthenticated mode: loopback alone does not identify the caller.
167    pub fn serve(&mut self, listener: TcpListener, credential: &str) -> std::io::Result<()> {
168        if credential.len() < 32 {
169            return Err(std::io::Error::new(
170                std::io::ErrorKind::InvalidInput,
171                "credential must be at least 32 characters",
172            ));
173        }
174        let port = listener.local_addr()?.port();
175        for conn in listener.incoming() {
176            let mut s = match conn {
177                Ok(s) => s,
178                Err(_) => continue,
179            };
180            let _ = s.set_read_timeout(Some(Duration::from_secs(5)));
181            let _ = s.set_write_timeout(Some(Duration::from_secs(5)));
182            let peer_ok = s.peer_addr().map(|a| a.ip().is_loopback()).unwrap_or(false);
183            let (code, out) = match read_request(&mut s) {
184                Some((head, body)) => match check_request(&head, peer_ok, port, credential) {
185                    Ok(()) => {
186                        let mut first = head.lines().next().unwrap_or("").split(' ');
187                        let (m, p) = (first.next().unwrap_or(""), first.next().unwrap_or("/"));
188                        self.handle(m, p, &body)
189                    }
190                    Err(e) => e,
191                },
192                None => err(400, "invalid argument", "malformed request"),
193            };
194            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());
195        }
196        Ok(())
197    }
198}
199
200const MAX_REQUEST: usize = 1 << 20;
201
202/// Read one HTTP request (head, body), bounded; `None` on malformed or oversized input.
203fn read_request(s: &mut std::net::TcpStream) -> Option<(String, String)> {
204    let mut buf: Vec<u8> = Vec::new();
205    let mut chunk = [0u8; 8192];
206    let (head_end, len) = loop {
207        let n = s.read(&mut chunk).ok()?;
208        if n == 0 {
209            return None;
210        }
211        buf.extend_from_slice(&chunk[..n]);
212        if buf.len() > MAX_REQUEST {
213            return None;
214        }
215        if let Some(i) = buf.windows(4).position(|w| w == b"\r\n\r\n") {
216            let head = String::from_utf8_lossy(&buf[..i]).to_string();
217            let cl = header(&head, "content-length")
218                .and_then(|v| v.parse::<usize>().ok())
219                .unwrap_or(0);
220            if cl > MAX_REQUEST {
221                return None;
222            }
223            break (i + 4, cl);
224        }
225    };
226    while buf.len() < head_end + len {
227        let n = s.read(&mut chunk).ok()?;
228        if n == 0 {
229            return None;
230        }
231        buf.extend_from_slice(&chunk[..n]);
232        if buf.len() > MAX_REQUEST + 8192 {
233            return None;
234        }
235    }
236    let head = String::from_utf8_lossy(&buf[..head_end - 4]).to_string();
237    let body = String::from_utf8_lossy(&buf[head_end..head_end + len]).to_string();
238    Some((head, body))
239}
240
241fn header(head: &str, name: &str) -> Option<String> {
242    head.lines().skip(1).find_map(|l| {
243        let (k, v) = l.split_once(':')?;
244        k.trim()
245            .eq_ignore_ascii_case(name)
246            .then(|| v.trim().to_string())
247    })
248}
249
250fn ct_eq(a: &str, b: &str) -> bool {
251    let (a, b) = (a.as_bytes(), b.as_bytes());
252    let mut d = u32::from(a.len() != b.len());
253    for i in 0..a.len().max(b.len()) {
254        d |= u32::from(a.get(i).copied().unwrap_or(0) ^ b.get(i).copied().unwrap_or(0));
255    }
256    d == 0
257}
258
259/// Admission for one request head. Pure so it is testable without a socket.
260pub fn check_request(
261    head: &str,
262    peer_is_loopback: bool,
263    port: u16,
264    credential: &str,
265) -> Result<(), (u16, String)> {
266    if !peer_is_loopback {
267        return Err(err(403, "unknown error", "peer is not loopback"));
268    }
269    let host_ok = header(head, "host")
270        .map(|h| {
271            h == format!("127.0.0.1:{port}")
272                || h == format!("localhost:{port}")
273                || h == format!("[::1]:{port}")
274        })
275        .unwrap_or(false);
276    if !host_ok {
277        return Err(err(403, "unknown error", "unexpected Host header"));
278    }
279    if header(head, "origin").is_some() {
280        return Err(err(
281            403,
282            "unknown error",
283            "browser-originated requests are refused",
284        ));
285    }
286    let presented = header(head, "authorization")
287        .and_then(|v| v.strip_prefix("Bearer ").map(str::to_string))
288        .unwrap_or_default();
289    if !ct_eq(&presented, credential) {
290        return Err(err(
291            401,
292            "unknown error",
293            "missing or wrong bridge credential",
294        ));
295    }
296    Ok(())
297}
298
299#[cfg(target_os = "macos")]
300pub mod macos {
301    use super::Backend;
302    use crate::mac::{self, Node};
303    use std::collections::HashMap;
304
305    /// Accessibility-resolved elements, background real-input actions.
306    pub struct AxBackend {
307        pub pid: i32,
308        els: HashMap<String, Node>,
309        n: u32,
310    }
311    impl AxBackend {
312        pub fn new(pid: i32) -> Self {
313            Self {
314                pid,
315                els: HashMap::new(),
316                n: 0,
317            }
318        }
319        fn get(&self, id: &str) -> Result<&Node, String> {
320            self.els
321                .get(id)
322                .ok_or_else(|| "unknown element".to_string())
323        }
324    }
325    impl Backend for AxBackend {
326        fn find(&mut self, using: &str, value: &str) -> Result<String, String> {
327            let exact = using == "link text" || using == "name";
328            let node = mac::ax_find(self.pid, value, exact)
329                .ok_or_else(|| format!("no element matching {value:?}"))?;
330            self.n += 1;
331            let id = format!("ax-{}", self.n);
332            self.els.insert(id.clone(), node);
333            Ok(id)
334        }
335        fn click(&mut self, id: &str) -> Result<(), String> {
336            let n = self.get(id)?.clone();
337            // AXPress first: works on inactive apps whose WKWebView swallows
338            // first-mouse clicks. Fall back to a real per-pid mouse click.
339            if mac::ax_press(&n) {
340                return Ok(());
341            }
342            mac::click_node(self.pid, &n)
343        }
344        fn send_keys(&mut self, id: &str, text: &str) -> Result<(), String> {
345            let n = self.get(id)?.clone();
346            mac::click_node(self.pid, &n)?; // focus inside the app, not the OS
347            mac::type_text(self.pid, text);
348            Ok(())
349        }
350        fn text(&mut self, id: &str) -> Result<String, String> {
351            Ok(mac::ax_text(self.get(id)?))
352        }
353        fn screenshot_png_b64(&mut self) -> Result<String, String> {
354            let w = mac::main_window(self.pid).ok_or("no window")?;
355            let p = std::env::temp_dir().join(format!("rkc-{}.png", std::process::id()));
356            if !mac::capture_window(w.id, &p.to_string_lossy()) {
357                return Err("screencapture failed".into());
358            }
359            let bytes = std::fs::read(&p).map_err(|e| e.to_string())?;
360            let _ = std::fs::remove_file(&p);
361            Ok(super::b64(&bytes))
362        }
363    }
364}
365
366#[cfg(test)]
367mod tests {
368    use super::*;
369    struct Fake(Vec<String>);
370    impl Backend for Fake {
371        fn find(&mut self, u: &str, v: &str) -> Result<String, String> {
372            self.0.push(format!("find {u} {v}"));
373            Ok("e1".into())
374        }
375        fn click(&mut self, id: &str) -> Result<(), String> {
376            self.0.push(format!("click {id}"));
377            Ok(())
378        }
379        fn send_keys(&mut self, _: &str, t: &str) -> Result<(), String> {
380            self.0.push(format!("keys {t}"));
381            Ok(())
382        }
383        fn text(&mut self, _: &str) -> Result<String, String> {
384            Ok("hi".into())
385        }
386        fn screenshot_png_b64(&mut self) -> Result<String, String> {
387            Ok(b64(b"abcd"))
388        }
389    }
390    #[test]
391    fn session_flow() {
392        let mut s = Server::new(Fake(vec![]));
393        let (_, r) = s.handle("POST", "/session", "{}");
394        assert!(r.contains("rkc-1"));
395        let (c, r) = s.handle(
396            "POST",
397            "/session/rkc-1/element",
398            r#"{"using":"link text","value":"Go"}"#,
399        );
400        assert_eq!(c, 200);
401        assert!(r.contains(ELEMENT_KEY));
402        assert_eq!(
403            s.handle("POST", "/session/rkc-1/element/e1/click", "{}").0,
404            200
405        );
406        assert_eq!(s.handle("POST", "/session/zzz/element", "{}").0, 404);
407        assert_eq!(s.backend.0, ["find link text Go", "click e1"]);
408        assert_eq!(b64(b"abcd"), "YWJjZA==");
409    }
410
411    #[test]
412    fn denied_click_never_reaches_backend() {
413        use crate::admission::{AdmissionDecision, SettleOutcome};
414        let gate = Arc::new(EffectGate::new(|r: &EffectRequest| {
415            if r.method == "click" {
416                AdmissionDecision::deny("clicks blocked")
417            } else {
418                AdmissionDecision::Allow
419            }
420        }));
421        let mut s = Server::new(Fake(vec![])).with_gate(gate.clone());
422        s.handle("POST", "/session", "{}");
423        s.handle(
424            "POST",
425            "/session/rkc-1/element",
426            r#"{"using":"css","value":"x"}"#,
427        );
428        let (c, r) = s.handle("POST", "/session/rkc-1/element/e1/click", "{}");
429        assert_eq!(c, 400);
430        assert!(r.contains("denied"), "{r}");
431        let (c, _) = s.handle(
432            "POST",
433            "/session/rkc-1/element/e1/value",
434            r#"{"text":"hi"}"#,
435        );
436        assert_eq!(c, 200);
437        assert_eq!(s.backend.0, ["find css x", "keys hi"]);
438        let outcomes: Vec<_> = gate.settlements().iter().map(|x| x.outcome).collect();
439        assert_eq!(outcomes, [SettleOutcome::Denied, SettleOutcome::Ok]);
440    }
441}