diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml new file mode 100644 index 0000000..9bec422 --- /dev/null +++ b/.github/workflows/ci.yml @@ -0,0 +1,34 @@ +name: CI + +on: + push: + pull_request: + +jobs: + rust: + name: Rust (${{ matrix.os }}) + runs-on: ${{ matrix.os }} + strategy: + fail-fast: false + matrix: + os: [ubuntu-latest, macos-latest, windows-latest] + steps: + - uses: actions/checkout@v4 + - uses: dtolnay/rust-toolchain@stable + with: + components: rustfmt, clippy + - run: cargo fmt --all -- --check + - run: cargo clippy --all-targets --all-features + - run: cargo test --all-targets --all-features + + extension: + name: Chrome extension syntax + runs-on: ubuntu-latest + steps: + - uses: actions/checkout@v4 + - uses: actions/setup-node@v4 + with: + node-version: 24 + - run: node --check assets/tmwd_cdp_bridge/background.js + - run: node --check assets/tmwd_cdp_bridge/content.js + - run: node --check assets/tmwd_cdp_bridge/popup.js diff --git a/Cargo.lock b/Cargo.lock index 27ce86b..7ae62cf 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -9,15 +9,18 @@ dependencies = [ "anyhow", "axum", "clap", + "dirs", + "dunce", "fs2", "html5ever", "libc", "markup5ever_rcdom", + "path-clean", "reqwest", "serde", "serde_json", "tokio", - "tower-http", + "url", "uuid", ] @@ -294,6 +297,27 @@ dependencies = [ "crypto-common", ] +[[package]] +name = "dirs" +version = "6.0.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c3e8aa94d75141228480295a7d0e7feb620b1a5ad9f12bc40be62411e38cce4e" +dependencies = [ + "dirs-sys", +] + +[[package]] +name = "dirs-sys" +version = "0.5.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e01a3366d27ee9890022452ee61b2b63a67e6f13f58900b651ff5665f0bb1fab" +dependencies = [ + "libc", + "option-ext", + "redox_users", + "windows-sys 0.61.2", +] + [[package]] name = "displaydoc" version = "0.2.5" @@ -305,6 +329,12 @@ dependencies = [ "syn", ] +[[package]] +name = "dunce" +version = "1.0.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "92773504d58c093f6de2459af4af33faa518c13451eb8f2b5698ed3d36e7c813" + [[package]] name = "equivalent" version = "1.0.2" @@ -764,6 +794,15 @@ version = "0.2.186" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "68ab91017fe16c622486840e4c83c9a37afeff978bd239b5293d61ece587de66" +[[package]] +name = "libredox" +version = "0.1.18" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c943259e342f1e06ff2da7a83eabdfe7f92ce10262688dbf1895ff0b3e6e4652" +dependencies = [ + "libc", +] + [[package]] name = "litemap" version = "0.8.2" @@ -870,6 +909,12 @@ version = "1.70.2" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "384b8ab6d37215f3c5301a95a4accb5d64aa607f1fcb26a11b5303878451b4fe" +[[package]] +name = "option-ext" +version = "0.2.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "04744f49eae99ab78e0d5c0b603ab218f515ea8cfe5a456d7629ad883a3b6e7d" + [[package]] name = "parking_lot" version = "0.12.5" @@ -893,6 +938,12 @@ dependencies = [ "windows-link", ] +[[package]] +name = "path-clean" +version = "1.0.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "17359afc20d7ab31fdb42bb844c8b3bb1dabd7dcf7e68428492da7f16966fcef" + [[package]] name = "percent-encoding" version = "2.3.2" @@ -1130,6 +1181,17 @@ dependencies = [ "bitflags", ] +[[package]] +name = "redox_users" +version = "0.5.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a4e608c6638b9c18977b00b475ac1f28d14e84b27d8d42f70e0bf1e3dec127ac" +dependencies = [ + "getrandom 0.2.17", + "libredox", + "thiserror 2.0.18", +] + [[package]] name = "reqwest" version = "0.12.28" diff --git a/Cargo.toml b/Cargo.toml index 5d4e265..fa3adf9 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -12,13 +12,16 @@ path = "src/main.rs" anyhow = "1.0" axum = { version = "0.7", features = ["json", "ws"] } clap = { version = "4.5", features = ["derive"] } +dirs = "6.0" +dunce = "1.0" fs2 = "0.4" html5ever = "0.27" libc = "0.2" markup5ever_rcdom = "0.3" +path-clean = "1.0" reqwest = { version = "0.12", default-features = false, features = ["blocking", "json", "rustls-tls"] } serde = { version = "1.0", features = ["derive"] } serde_json = "1.0" tokio = { version = "1.48", features = ["macros", "process", "rt-multi-thread", "signal", "sync", "time"] } -tower-http = { version = "0.6", features = ["cors"] } +url = "2.5" uuid = { version = "1.18", features = ["v4"] } diff --git a/assets/tmwd_cdp_bridge/background.js b/assets/tmwd_cdp_bridge/background.js index 0009ace..baeee07 100644 --- a/assets/tmwd_cdp_bridge/background.js +++ b/assets/tmwd_cdp_bridge/background.js @@ -20,6 +20,7 @@ const DEFAULT_WS_PORT = 18765; const CLI_API_PORT = 18767; let wsPort = DEFAULT_WS_PORT; const browserId = `browser-${crypto.randomUUID()}`; +const extensionVersion = chrome.runtime.getManifest().version; let profileId = null; let profileLabel = null; @@ -40,7 +41,7 @@ function withClientIdentity(payload) { } async function handleExtMessage(msg, sender) { - if (msg.cmd === 'status') return handleStatus(); + if (msg.cmd === 'status') return await handleStatus(); if (msg.cmd === 'setPort') return await handleSetPort(msg); if (msg.cmd === 'setProfileLabel') return await handleSetProfileLabel(msg); lastCommandAt = Date.now(); @@ -67,7 +68,7 @@ async function handleExtMessage(msg, sender) { if (msg.allowFocus === true && tab.windowId) await chrome.windows.update(tab.windowId, { focused: true }); return { ok: true }; } else { - const tabs = (await chrome.tabs.query({})).filter(t => isScriptable(t.url)); + const tabs = await queryScriptableTabs(); const data = tabs.map(t => ({ id: t.id, url: t.url, title: t.title, active: t.active, windowId: t.windowId })); return { ok: true, data }; } @@ -109,7 +110,7 @@ async function handleExtMessage(msg, sender) { return { ok: false, error: 'Unknown cmd: ' + msg.cmd }; } -function handleStatus() { +async function handleStatus() { return { ok: true, data: { @@ -119,7 +120,10 @@ function handleStatus() { lastCommandAt, browserId, profileId, - profileLabel + profileLabel, + extensionId: chrome.runtime.id, + extensionVersion, + fileSchemeAccess: await chrome.extension.isAllowedFileSchemeAccess() } }; } @@ -645,7 +649,7 @@ async function handleBatch(msg, sender) { if (c.cmd === 'cookies') { R.push(await handleCookies(c, sender)); } else if (c.cmd === 'tabs') { - const tabs = (await chrome.tabs.query({})).filter(t => isScriptable(t.url)); + const tabs = await queryScriptableTabs(); R.push({ ok: true, data: tabs.map(t => ({ id: t.id, url: t.url, title: t.title, active: t.active, windowId: t.windowId })) }); } else if (c.cmd === 'cdp') { const tabId = c.tabId || msg.tabId || sender.tab?.id; @@ -687,11 +691,16 @@ async function handleCDP(msg, sender) { return { ok: false, error: e.message }; } } -// Filter out chrome:// and other internal tabs that can't be scripted -const isScriptable = url => url && /^https?:/.test(url); +// Filter out chrome:// and other internal tabs that can't be scripted. +const isScriptable = (url, fileSchemeAccess = false) => !!url && (/^https?:/i.test(url) || (fileSchemeAccess && /^file:/i.test(url))); + +async function queryScriptableTabs() { + const fileSchemeAccess = await chrome.extension.isAllowedFileSchemeAccess(); + return (await chrome.tabs.query({})).filter(tab => isScriptable(tab.url, fileSchemeAccess)); +} async function injectContentScriptsIntoExistingTabs() { - const tabs = (await chrome.tabs.query({})).filter(t => isScriptable(t.url)); + const tabs = await queryScriptableTabs(); for (const tab of tabs) { try { await chrome.scripting.executeScript({ @@ -1029,9 +1038,11 @@ async function connectWS() { console.log('[TMWD-WS] Connected!'); scheduleKeepalive(); // Keep SW alive while connected await loadClientIdentity(); - const tabs = (await chrome.tabs.query({})).filter(t => isScriptable(t.url)); + const tabs = await queryScriptableTabs(); ws.send(JSON.stringify(withClientIdentity({ type: 'ext_ready', + extension_version: extensionVersion, + file_scheme_access: await chrome.extension.isAllowedFileSchemeAccess(), tabs: tabs.map(t => ({ id: t.id, url: t.url, title: t.title })) }))); console.log('[TMWD-WS] Sent ext_ready with', tabs.length, 'tabs'); @@ -1103,10 +1114,12 @@ chrome.runtime.onInstalled.addListener(() => { // Sync tab list on changes async function sendTabsUpdate() { if (!ws || ws.readyState !== WebSocket.OPEN) return; - const tabs = (await chrome.tabs.query({})).filter(t => isScriptable(t.url) && !/streamlit/i.test(t.title)); + const tabs = (await queryScriptableTabs()).filter(t => !/streamlit/i.test(t.title)); await loadClientIdentity(); ws.send(JSON.stringify(withClientIdentity({ type: 'tabs_update', + extension_version: extensionVersion, + file_scheme_access: await chrome.extension.isAllowedFileSchemeAccess(), tabs: tabs.map(t => ({ id: t.id, url: t.url, title: t.title })) }))); } diff --git a/assets/tmwd_cdp_bridge/manifest.json b/assets/tmwd_cdp_bridge/manifest.json index 911e21f..40cf8ca 100644 --- a/assets/tmwd_cdp_bridge/manifest.json +++ b/assets/tmwd_cdp_bridge/manifest.json @@ -1,7 +1,7 @@ { "manifest_version": 3, "name": "Agent Browser CLI Bridge", - "version": "2.0", + "version": "2.1", "description": "Browser control bridge for agent-browser-cli", "permissions": [ "cookies", diff --git a/assets/tmwd_cdp_bridge/popup.html b/assets/tmwd_cdp_bridge/popup.html index 1a4fb2f..fb2974b 100644 --- a/assets/tmwd_cdp_bridge/popup.html +++ b/assets/tmwd_cdp_bridge/popup.html @@ -21,6 +21,7 @@

🔌 Bridge

读取连接状态...
+
读取文件网址权限...
diff --git a/assets/tmwd_cdp_bridge/popup.js b/assets/tmwd_cdp_bridge/popup.js index 5a1a131..ae46742 100644 --- a/assets/tmwd_cdp_bridge/popup.js +++ b/assets/tmwd_cdp_bridge/popup.js @@ -20,10 +20,16 @@ async function refreshBridgeStatus() { const data = resp.data || {}; portInput.value = data.wsPort || 18765; status.textContent = `状态: ${data.wsConnected ? '已连接' : '未连接'} ${data.wsUrl || ''}`; + const fileAccessStatus = document.getElementById('fileAccessStatus'); + fileAccessStatus.textContent = `文件网址访问: ${data.fileSchemeAccess ? '已启用' : '未启用'}`; + fileAccessStatus.className = data.fileSchemeAccess ? 'status' : 'error'; renderProfileStatus(data); } catch (e) { status.textContent = '状态读取失败: ' + e.message; status.className = 'error'; + const fileAccessStatus = document.getElementById('fileAccessStatus'); + fileAccessStatus.textContent = '文件网址权限读取失败: ' + e.message; + fileAccessStatus.className = 'error'; const profileStatus = document.getElementById('profileStatus'); profileStatus.textContent = 'Profile 读取失败: ' + e.message; profileStatus.className = 'error'; diff --git a/skills/agent-browser-cli/SKILL.md b/skills/agent-browser-cli/SKILL.md index 16f5a59..d8823fc 100644 --- a/skills/agent-browser-cli/SKILL.md +++ b/skills/agent-browser-cli/SKILL.md @@ -30,6 +30,30 @@ agent-browser-cli logs --tail 100 补充一个容易误判的点:`status` / `doctor` 里看到 `daemon_not_running`,如果此时还没有执行 `tabs` / `open` / `exec` / `scan`,通常只是 daemon 按需常驻而未启动,不代表故障。只有目标命令已经失败,或者日志/输出明确提示端口、扩展、标签页异常时,才进入排障。 +## 打开本地文件 + +`open` 支持 HTTP(S) URL、绝对 `file://` URL 和本机文件路径。相对路径以执行 CLI 时的当前目录为基准;支持 `~`、中文、空格和符号链接,符号链接会解析到真实文件。所有普通文件类型均可打开,目录会被拒绝。 + +```bash +agent-browser-cli open ./demo.html +agent-browser-cli open ~/Documents/report.pdf +agent-browser-cli open 'file:///Users/me/My%20Files/demo.html#intro' +agent-browser-cli open ./demo.html --timeout 15 +``` + +本地文件要求 Chrome 扩展版本至少为 2.1,并在 `chrome://extensions` → Agent Browser CLI Bridge → 详情中开启“允许访问文件网址”。权限关闭时 `open` 会在创建标签前失败;`status` / `doctor` 会按 Profile 报告 `file_scheme_access`,但该可选能力不影响普通 HTTP(S) 健康状态。 + +输入分类规则: + +- `./`、`../`、`~/`、绝对路径和显式 `file://` 表示本地文件;文件必须存在且是普通文件。 +- 裸输入若对应当前目录中已存在的文件,文件优先;否则按 Web 地址处理。要强制打开网站,显式写 `https://`。 +- 普通路径中的 `?` 和 `#` 是文件名字符;仅显式 `file://` URL 使用 query/fragment 语义。 +- `localhost`、回环/私网 IP、单标签主机和 `.local` 的无协议地址默认使用 HTTP;其他域名默认 HTTPS;端口 80/443 分别强制 HTTP/HTTPS。 +- 只支持 `http:`、`https:` 和 `file:`;其他显式 scheme 会返回 `unsupported_scheme`。 +- 异平台路径会明确报错,例如 macOS/Linux 不会把 `C:\\Users\\...` 错当作网址。 + +本地文件 `open` 默认最多等待 10 秒,直到新标签成为可控制会话;可用 `--timeout` 调整。超时后标签会保留,错误结果包含稳定的 `error_code` 和 `opened_tab_id`。活动 `open` 成功后新标签成为默认会话,`--background` 不切换默认会话。 + ## 常用命令优先级 先区分三个入口: diff --git a/src/cli.rs b/src/cli.rs index b10ad0b..08dc6e2 100644 --- a/src/cli.rs +++ b/src/cli.rs @@ -341,7 +341,7 @@ struct ConsoleListArgs { #[derive(Debug, Args)] struct OpenArgs { - url: String, + target: String, #[arg(long)] background: bool, #[arg(long)] @@ -358,7 +358,7 @@ struct OpenArgs { session: Option, #[arg(long = "group-title")] group_title: Option, - #[arg(long, default_value_t = 30.0)] + #[arg(long, default_value_t = 10.0)] timeout: f64, } @@ -578,12 +578,19 @@ pub fn run() -> Result<()> { CommandKind::Network(args) => run_network_command(args), CommandKind::Console(args) => run_console_command(args), CommandKind::Open(args) => { + if !args.timeout.is_finite() || args.timeout <= 0.0 { + return Err(anyhow!("open --timeout 必须是大于 0 的有限秒数")); + } ensure_server()?; + let cwd = env::current_dir().context("无法读取当前工作目录")?; + let request_timeout = (args.timeout + 20.0).max(30.0); print_json(request( "POST", "/open", Some(json!({ - "url": args.url, + "target": args.target.clone(), + "url": args.target, + "cwd": cwd, "active": !args.background, "window": args.window, "allow_focus": args.focus, @@ -592,8 +599,9 @@ pub fn run() -> Result<()> { "profile": args.profile, "session": args.session, "group_title": args.group_title, + "readiness_timeout": args.timeout, })), - args.timeout, + request_timeout, )?); Ok(()) } @@ -668,10 +676,17 @@ fn request(method: &str, path: &str, payload: Option, timeout_secs: f64) .timeout(Duration::from_secs_f64(timeout_secs.max(0.1))) .build()?; let url = format!("http://{HOST}:{PORT}{path}"); + let token = config::load_api_token()?; + let with_token = |request: reqwest::blocking::RequestBuilder| { + if let Some(token) = token.as_deref() { + request.header(config::API_TOKEN_HEADER, token) + } else { + request + } + }; let response = match method { - "GET" => client.get(url).send()?, - "POST" => client - .post(url) + "GET" => with_token(client.get(url)).send()?, + "POST" => with_token(client.post(url)) .json(&payload.unwrap_or_else(|| json!({}))) .send()?, _ => return Err(anyhow!("不支持的 HTTP 方法: {method}")), @@ -693,6 +708,7 @@ fn ensure_server() -> Result<()> { let lock_path = project_dir().join(".agent-browser-cli.lock"); let lock = OpenOptions::new() .create(true) + .truncate(false) .read(true) .write(true) .open(lock_path)?; @@ -909,6 +925,11 @@ fn doctor_value() -> Result { .and_then(|v| v.pointer("/connection/active_tabs")) .and_then(Value::as_u64) .unwrap_or(0); + let file_scheme_profiles = health + .as_ref() + .and_then(|v| v.pointer("/connection/profiles")) + .cloned() + .unwrap_or_else(|| json!([])); let running = health .as_ref() .and_then(|v| v.get("running")) @@ -974,7 +995,14 @@ fn doctor_value() -> Result { "name": "active_tabs", "ok": active_tabs > 0, "active_tabs": active_tabs, - "hint": if active_tabs > 0 { Value::Null } else { json!("Chrome 需要至少打开一个普通 http/https 网页标签页") } + "hint": if active_tabs > 0 { Value::Null } else { json!("Chrome 需要至少打开一个普通 http/https/file 网页标签页") } + })); + checks.push(json!({ + "name": "file_scheme_access", + "ok": true, + "optional": true, + "profiles": file_scheme_profiles, + "hint": "文件网址访问是可选能力;关闭时不影响普通 http/https 页面控制" })); let ok = checks diff --git a/src/config.rs b/src/config.rs index 5fe911e..9751259 100644 --- a/src/config.rs +++ b/src/config.rs @@ -1,11 +1,13 @@ use anyhow::{anyhow, Context, Result}; use serde::{Deserialize, Serialize}; use std::env; -use std::fs; +use std::fs::{self, OpenOptions}; +use std::io::{ErrorKind, Write}; use std::path::PathBuf; pub const DEFAULT_EXTENSION_PORT: u16 = 18765; pub const CLI_API_PORT: u16 = 18767; +pub const API_TOKEN_HEADER: &str = "x-agent-browser-token"; #[derive(Debug, Clone, Serialize, Deserialize)] pub struct AppConfig { @@ -94,6 +96,70 @@ pub fn daemon_log_path() -> Result { Ok(log_dir()?.join("daemon.log")) } +pub fn api_token_path() -> Result { + Ok(user_config_dir()?.join("api-token")) +} + +pub fn load_api_token() -> Result> { + let path = api_token_path()?; + let content = match fs::read_to_string(&path) { + Ok(content) => content, + Err(err) if err.kind() == ErrorKind::NotFound => return Ok(None), + Err(err) => { + return Err(err).with_context(|| format!("读取 API token 失败: {}", path.display())) + } + }; + let token = content.trim(); + if token.is_empty() { + return Err(anyhow!("API token 文件为空: {}", path.display())); + } + Ok(Some(token.to_string())) +} + +pub fn save_api_token(token: &str) -> Result<()> { + if token.is_empty() { + return Err(anyhow!("API token 不能为空")); + } + let path = api_token_path()?; + let parent = path + .parent() + .ok_or_else(|| anyhow!("API token 路径缺少父目录: {}", path.display()))?; + fs::create_dir_all(parent) + .with_context(|| format!("创建配置目录失败: {}", parent.display()))?; + let temp_path = parent.join(format!(".api-token-{}.tmp", std::process::id())); + let mut options = OpenOptions::new(); + options.create(true).truncate(true).write(true); + #[cfg(unix)] + { + use std::os::unix::fs::OpenOptionsExt; + options.mode(0o600); + } + let result = (|| -> Result<()> { + let mut file = options + .open(&temp_path) + .with_context(|| format!("创建 API token 临时文件失败: {}", temp_path.display()))?; + file.write_all(token.as_bytes())?; + file.write_all(b"\n")?; + file.sync_all()?; + #[cfg(unix)] + { + use std::os::unix::fs::PermissionsExt; + file.set_permissions(fs::Permissions::from_mode(0o600))?; + } + #[cfg(windows)] + if path.exists() { + fs::remove_file(&path)?; + } + fs::rename(&temp_path, &path) + .with_context(|| format!("保存 API token 失败: {}", path.display()))?; + Ok(()) + })(); + if result.is_err() { + let _ = fs::remove_file(&temp_path); + } + result +} + pub fn ensure_log_file() -> Result { let path = daemon_log_path()?; if let Some(parent) = path.parent() { diff --git a/src/main.rs b/src/main.rs index 4041273..00e2798 100644 --- a/src/main.rs +++ b/src/main.rs @@ -1,6 +1,7 @@ mod cli; mod config; mod html; +mod open_target; mod protocol; mod server; diff --git a/src/open_target.rs b/src/open_target.rs new file mode 100644 index 0000000..8338fe2 --- /dev/null +++ b/src/open_target.rs @@ -0,0 +1,593 @@ +use path_clean::PathClean; +use serde_json::{json, Value}; +use std::fmt; +use std::fs; +use std::path::{Path, PathBuf}; +use url::{Host, Url}; + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum TargetKind { + Web, + LocalResource, +} + +#[derive(Debug, Clone)] +pub struct ResolvedOpenTarget { + pub navigation_url: String, + pub kind: TargetKind, +} + +#[derive(Debug, Clone)] +pub struct OpenTargetError { + pub code: &'static str, + pub message: String, + pub details: Value, +} + +impl OpenTargetError { + fn new(code: &'static str, message: impl Into, details: Value) -> Self { + Self { + code, + message: message.into(), + details, + } + } +} + +impl fmt::Display for OpenTargetError { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + f.write_str(&self.message) + } +} + +impl std::error::Error for OpenTargetError {} + +pub fn resolve_open_target( + input: &str, + cwd: Option<&Path>, +) -> Result { + if input.is_empty() || input.trim().is_empty() { + return Err(OpenTargetError::new( + "invalid_open_target", + "Open Target must not be empty", + json!({ "open_target": input }), + )); + } + if let Some(cwd) = cwd { + if !cwd.is_absolute() || !cwd.is_dir() { + return Err(OpenTargetError::new( + "invalid_cwd", + "cwd must be an absolute existing directory", + json!({ "open_target": input, "cwd": cwd }), + )); + } + } + + if let Some(platform) = foreign_path_platform(input) { + return Err(OpenTargetError::new( + "foreign_platform_path", + format!( + "The Open Target looks like a {platform} path, but the CLI is running on {}", + std::env::consts::OS + ), + json!({ + "open_target": input, + "host_platform": std::env::consts::OS, + "detected_platform": platform, + }), + )); + } + + if cfg!(windows) && is_explicit_path(input) { + return resolve_path_input(input, cwd); + } + + if let Some(scheme) = explicit_scheme(input) { + return match scheme.as_str() { + "http" | "https" => resolve_explicit_web(input, &scheme), + "file" => resolve_file_url(input), + _ if looks_like_host_port(input) => resolve_implicit_web(input), + _ => Err(OpenTargetError::new( + "unsupported_scheme", + format!("Unsupported Open Target scheme: {scheme}"), + json!({ "scheme": scheme, "supported_schemes": ["http", "https", "file"] }), + )), + }; + } + + if is_explicit_path(input) { + return resolve_path_input(input, cwd); + } + + if let Some(cwd) = cwd { + if cwd.join(input).exists() { + return resolve_path_input(input, Some(cwd)); + } + } + + resolve_implicit_web(input) +} + +fn resolve_explicit_web( + input: &str, + expected_scheme: &str, +) -> Result { + let parsed = Url::parse(input).map_err(|err| { + OpenTargetError::new( + "invalid_http_url", + format!("Invalid {} URL: {err}", expected_scheme.to_uppercase()), + json!({ "open_target": input, "scheme": expected_scheme }), + ) + })?; + if parsed.host().is_none() { + return Err(OpenTargetError::new( + "invalid_http_url", + format!( + "Invalid {} URL: host is missing", + expected_scheme.to_uppercase() + ), + json!({ "open_target": input, "scheme": expected_scheme }), + )); + } + Ok(ResolvedOpenTarget { + navigation_url: input.to_string(), + kind: TargetKind::Web, + }) +} + +fn resolve_implicit_web(input: &str) -> Result { + let probe = Url::parse(&format!("https://{input}")).map_err(|err| { + OpenTargetError::new( + "invalid_open_target", + format!("Open Target is neither a valid local path nor a valid web address: {err}"), + json!({ "open_target": input }), + ) + })?; + let host = probe.host().ok_or_else(|| { + OpenTargetError::new( + "invalid_open_target", + "Open Target is missing a web host", + json!({ "open_target": input }), + ) + })?; + let scheme = match explicit_port(input) { + Some(80) => "http", + Some(443) => "https", + _ if host_defaults_to_http(host) => "http", + _ => "https", + }; + Ok(ResolvedOpenTarget { + navigation_url: format!("{scheme}://{input}"), + kind: TargetKind::Web, + }) +} + +fn host_defaults_to_http(host: Host<&str>) -> bool { + match host { + Host::Domain(domain) => { + let lower = domain.to_ascii_lowercase(); + lower == "localhost" + || lower.ends_with(".localhost") + || lower.ends_with(".local") + || !lower.contains('.') + } + Host::Ipv4(ip) => { + ip.is_loopback() || ip.is_private() || ip.is_link_local() || ip.is_unspecified() + } + Host::Ipv6(ip) => { + ip.is_loopback() + || ip.is_unique_local() + || ip.is_unicast_link_local() + || ip.is_unspecified() + } + } +} + +fn resolve_file_url(input: &str) -> Result { + if !input.to_ascii_lowercase().starts_with("file://") { + return Err(OpenTargetError::new( + "invalid_file_url", + "Relative file URLs are not supported; pass a relative filesystem path instead", + json!({ "open_target": input }), + )); + } + let mut parsed = Url::parse(input).map_err(|err| { + OpenTargetError::new( + "invalid_file_url", + format!("Invalid file URL: {err}"), + json!({ "open_target": input }), + ) + })?; + if !parsed.username().is_empty() || parsed.password().is_some() || parsed.port().is_some() { + return Err(OpenTargetError::new( + "invalid_file_url", + "File URL must not contain credentials or a port", + json!({ "open_target": input }), + )); + } + + let host = parsed.host_str().unwrap_or_default().to_string(); + #[cfg(not(windows))] + if !host.is_empty() && !host.eq_ignore_ascii_case("localhost") { + return Err(OpenTargetError::new( + "unsupported_file_host", + format!("File URL host is not supported on this platform: {host}"), + json!({ "open_target": input, "host": host, "host_platform": std::env::consts::OS }), + )); + } + if host.eq_ignore_ascii_case("localhost") { + parsed.set_host(None).map_err(|_| { + OpenTargetError::new( + "invalid_file_url", + "Could not normalize localhost file URL", + json!({ "open_target": input }), + ) + })?; + } + + let query = parsed.query().map(str::to_string); + let fragment = parsed.fragment().map(str::to_string); + parsed.set_query(None); + parsed.set_fragment(None); + let path = parsed.to_file_path().map_err(|_| { + OpenTargetError::new( + "invalid_file_url", + "File URL must be absolute and map to a path on this platform", + json!({ "open_target": input }), + ) + })?; + if !path.is_absolute() { + return Err(OpenTargetError::new( + "invalid_file_url", + "Relative file URLs are not supported; pass a relative filesystem path instead", + json!({ "open_target": input }), + )); + } + resolve_local_path(path, input, query.as_deref(), fragment.as_deref()) +} + +fn resolve_path_input( + input: &str, + cwd: Option<&Path>, +) -> Result { + let path = expand_current_user_home(input)?; + let absolute = if path.is_absolute() { + path + } else { + let cwd = validate_cwd(cwd, input)?; + cwd.join(path) + } + .clean(); + resolve_local_path(absolute, input, None, None) +} + +fn validate_cwd<'a>(cwd: Option<&'a Path>, input: &str) -> Result<&'a Path, OpenTargetError> { + let cwd = cwd.ok_or_else(|| { + OpenTargetError::new( + "missing_cwd", + "A working directory is required for a relative Open Target", + json!({ "open_target": input }), + ) + })?; + if !cwd.is_absolute() || !cwd.is_dir() { + return Err(OpenTargetError::new( + "invalid_cwd", + "cwd must be an absolute existing directory", + json!({ "open_target": input, "cwd": cwd }), + )); + } + Ok(cwd) +} + +fn resolve_local_path( + path: PathBuf, + input: &str, + query: Option<&str>, + fragment: Option<&str>, +) -> Result { + if !path.exists() { + return Err(OpenTargetError::new( + "local_resource_not_found", + format!("Local Resource does not exist: {}", path.display()), + json!({ "open_target": input, "resolved_path": path }), + )); + } + let canonical = dunce::canonicalize(&path).map_err(|err| { + OpenTargetError::new( + "local_resource_not_found", + format!("Could not resolve Local Resource {}: {err}", path.display()), + json!({ "open_target": input, "resolved_path": path }), + ) + })?; + if canonical.to_str().is_none() { + return Err(OpenTargetError::new( + "non_utf8_local_resource", + "Local Resource path cannot be represented as UTF-8", + json!({ "open_target": input }), + )); + } + let metadata = fs::metadata(&canonical).map_err(|err| { + OpenTargetError::new( + "local_resource_not_found", + format!( + "Could not inspect Local Resource {}: {err}", + canonical.display() + ), + json!({ "open_target": input, "resolved_path": canonical }), + ) + })?; + if !metadata.is_file() { + return Err(OpenTargetError::new( + "local_resource_not_file", + format!( + "Local Resource must be a regular file: {}", + canonical.display() + ), + json!({ + "open_target": input, + "resolved_path": canonical, + "actual_type": if metadata.is_dir() { "directory" } else { "non_regular" }, + }), + )); + } + let mut url = Url::from_file_path(&canonical).map_err(|_| { + OpenTargetError::new( + "invalid_file_url", + "Could not convert Local Resource to a file URL", + json!({ "open_target": input, "resolved_path": canonical }), + ) + })?; + url.set_query(query); + url.set_fragment(fragment); + Ok(ResolvedOpenTarget { + navigation_url: url.to_string(), + kind: TargetKind::LocalResource, + }) +} + +fn expand_current_user_home(input: &str) -> Result { + if input == "~" || input.starts_with("~/") || (cfg!(windows) && input.starts_with("~\\")) { + let home = dirs::home_dir().ok_or_else(|| { + OpenTargetError::new( + "home_directory_unavailable", + "Could not determine the current user's home directory", + json!({ "open_target": input }), + ) + })?; + if input == "~" { + return Ok(home); + } + return Ok(home.join(&input[2..])); + } + Ok(PathBuf::from(input)) +} + +fn explicit_scheme(input: &str) -> Option { + let (scheme, _) = input.split_once(':')?; + let mut chars = scheme.chars(); + if !chars.next()?.is_ascii_alphabetic() + || !chars.all(|c| c.is_ascii_alphanumeric() || matches!(c, '+' | '-' | '.')) + { + return None; + } + Some(scheme.to_ascii_lowercase()) +} + +fn looks_like_host_port(input: &str) -> bool { + explicit_port(input).is_some() +} + +fn explicit_port(input: &str) -> Option { + let authority = input.split(['/', '?', '#']).next()?; + if authority.starts_with('[') { + let end = authority.find(']')?; + return authority.get(end + 1..)?.strip_prefix(':')?.parse().ok(); + } + let (host, port) = authority.rsplit_once(':')?; + (!host.is_empty()).then(|| port.parse().ok()).flatten() +} + +fn is_explicit_path(input: &str) -> bool { + input == "~" + || input.starts_with("~/") + || input.starts_with("./") + || input.starts_with("../") + || input.starts_with(".\\") + || input.starts_with("..\\") + || Path::new(input).is_absolute() +} + +#[cfg(not(windows))] +fn foreign_path_platform(input: &str) -> Option<&'static str> { + let bytes = input.as_bytes(); + let drive = bytes.len() >= 3 + && bytes[0].is_ascii_alphabetic() + && bytes[1] == b':' + && matches!(bytes[2], b'\\' | b'/'); + let unc = input.starts_with("\\\\"); + let relative = input.starts_with(".\\") || input.starts_with("..\\"); + let home = input.starts_with("~\\"); + (drive || unc || relative || home).then_some("windows") +} + +#[cfg(windows)] +fn foreign_path_platform(input: &str) -> Option<&'static str> { + (input.starts_with("/Users/") || input.starts_with("/home/") || input.starts_with("/tmp/")) + .then_some("unix") +} + +#[cfg(test)] +mod tests { + use super::*; + use std::fs; + + fn temp_dir() -> PathBuf { + let dir = std::env::temp_dir().join(format!( + "agent-browser-open-target-{}", + uuid::Uuid::new_v4() + )); + fs::create_dir_all(&dir).unwrap(); + dir + } + + #[test] + fn existing_bare_path_wins_over_web() { + let dir = temp_dir(); + fs::write(dir.join("example.com"), "ok").unwrap(); + let resolved = resolve_open_target("example.com", Some(&dir)).unwrap(); + assert_eq!(resolved.kind, TargetKind::LocalResource); + assert!(resolved.navigation_url.starts_with("file://")); + fs::remove_dir_all(dir).unwrap(); + } + + #[test] + fn missing_bare_file_shape_falls_back_to_web() { + let dir = temp_dir(); + let resolved = resolve_open_target("demo.html", Some(&dir)).unwrap(); + assert_eq!(resolved.navigation_url, "https://demo.html"); + fs::remove_dir_all(dir).unwrap(); + } + + #[test] + fn explicit_missing_path_is_an_error() { + let dir = temp_dir(); + let err = resolve_open_target("./missing.html", Some(&dir)).unwrap_err(); + assert_eq!(err.code, "local_resource_not_found"); + fs::remove_dir_all(dir).unwrap(); + } + + #[test] + fn directories_are_rejected() { + let dir = temp_dir(); + let err = resolve_open_target(dir.to_str().unwrap(), None).unwrap_err(); + assert_eq!(err.code, "local_resource_not_file"); + fs::remove_dir_all(dir).unwrap(); + } + + #[test] + fn explicit_web_url_is_preserved() { + let input = "HTTPS://EXAMPLE.com:443/a/../b?q=a%2Fb"; + let resolved = resolve_open_target(input, None).unwrap(); + assert_eq!(resolved.navigation_url, input); + } + + #[test] + fn local_hosts_default_to_http_and_public_hosts_to_https() { + assert_eq!( + resolve_open_target("localhost:3000", None) + .unwrap() + .navigation_url, + "http://localhost:3000" + ); + assert_eq!( + resolve_open_target("printer.local", None) + .unwrap() + .navigation_url, + "http://printer.local" + ); + assert_eq!( + resolve_open_target("example.com", None) + .unwrap() + .navigation_url, + "https://example.com" + ); + } + + #[test] + fn standard_ports_override_host_defaults() { + assert_eq!( + resolve_open_target("example.com:80", None) + .unwrap() + .navigation_url, + "http://example.com:80" + ); + assert_eq!( + resolve_open_target("localhost:443", None) + .unwrap() + .navigation_url, + "https://localhost:443" + ); + } + + #[test] + fn unsupported_schemes_are_rejected() { + let err = resolve_open_target("mailto:user@example.com", None).unwrap_err(); + assert_eq!(err.code, "unsupported_scheme"); + } + + #[cfg(unix)] + #[test] + fn path_hash_and_question_mark_are_filename_characters() { + let dir = temp_dir(); + let filename = "report#1?draft.html"; + fs::write(dir.join(filename), "ok").unwrap(); + let resolved = resolve_open_target(&format!("./{filename}"), Some(&dir)).unwrap(); + assert!(resolved.navigation_url.contains("report%231%3Fdraft.html")); + fs::remove_dir_all(dir).unwrap(); + } + + #[test] + fn relative_file_urls_are_rejected() { + let err = resolve_open_target("file:demo.html", None).unwrap_err(); + assert_eq!(err.code, "invalid_file_url"); + } + + #[test] + fn relative_paths_require_a_valid_absolute_cwd() { + let err = resolve_open_target("./demo.html", None).unwrap_err(); + assert_eq!(err.code, "missing_cwd"); + let err = resolve_open_target("./demo.html", Some(Path::new("relative"))).unwrap_err(); + assert_eq!(err.code, "invalid_cwd"); + let err = + resolve_open_target("https://example.com", Some(Path::new("relative"))).unwrap_err(); + assert_eq!(err.code, "invalid_cwd"); + } + + #[cfg(unix)] + #[test] + fn symbolic_links_resolve_to_the_real_file() { + use std::os::unix::fs::symlink; + + let dir = temp_dir(); + let real = dir.join("real.html"); + let link = dir.join("current.html"); + fs::write(&real, "ok").unwrap(); + symlink(&real, &link).unwrap(); + let resolved = resolve_open_target(link.to_str().unwrap(), None).unwrap(); + assert_eq!( + Url::parse(&resolved.navigation_url) + .unwrap() + .to_file_path() + .unwrap(), + dunce::canonicalize(real).unwrap() + ); + fs::remove_dir_all(dir).unwrap(); + } + + #[test] + fn file_url_preserves_query_and_fragment() { + let dir = temp_dir(); + let file = dir.join("demo.html"); + fs::write(&file, "ok").unwrap(); + let mut url = Url::from_file_path(&file).unwrap(); + url.set_query(Some("theme=dark")); + url.set_fragment(Some("intro")); + let resolved = resolve_open_target(url.as_str(), None).unwrap(); + assert!(resolved.navigation_url.ends_with("?theme=dark#intro")); + fs::remove_dir_all(dir).unwrap(); + } + + #[cfg(not(windows))] + #[test] + fn windows_paths_are_rejected_on_unix() { + for input in [ + r"C:\Users\Alice\demo.html", + r"\\server\share\demo.html", + r".\demo.html", + r"~\demo.html", + ] { + let err = resolve_open_target(input, None).unwrap_err(); + assert_eq!(err.code, "foreign_platform_path", "input: {input}"); + } + } +} diff --git a/src/protocol.rs b/src/protocol.rs index f23df2c..6a57370 100644 --- a/src/protocol.rs +++ b/src/protocol.rs @@ -38,6 +38,8 @@ pub struct Session { pub browser_id: String, pub profile_id: String, pub profile_label: Option, + pub extension_version: Option, + pub file_scheme_access: Option, pub info: TabInfo, pub sender: mpsc::UnboundedSender, pub disconnected_at: Option, @@ -134,6 +136,7 @@ pub struct DriverState { pub pending: HashMap, pub default_session_key: Option, pub latest_session_key: Option, + pub preferred_default_session_key: Option, pub active_exec_sessions: HashMap, pub acked: HashSet, } @@ -149,6 +152,10 @@ pub enum WsIncoming { profile_id: String, #[serde(default)] profile_label: Option, + #[serde(default)] + extension_version: Option, + #[serde(default)] + file_scheme_access: Option, tabs: Vec, }, #[serde(rename = "tabs_update")] @@ -159,6 +166,10 @@ pub enum WsIncoming { profile_id: String, #[serde(default)] profile_label: Option, + #[serde(default)] + extension_version: Option, + #[serde(default)] + file_scheme_access: Option, tabs: Vec, }, #[serde(rename = "ack")] diff --git a/src/server.rs b/src/server.rs index 51aabd2..c2c449d 100644 --- a/src/server.rs +++ b/src/server.rs @@ -1,13 +1,16 @@ +use crate::open_target::{resolve_open_target, OpenTargetError, TargetKind}; use crate::protocol::{ DriverState, ElementDomInfo, ElementRef, ExecResult, RectInfo, Session, SnapshotCache, TabInfo, WsIncoming, }; use crate::{config, html}; use anyhow::{anyhow, Result}; +use axum::body::Body; use axum::extract::ws::{Message, WebSocket, WebSocketUpgrade}; use axum::extract::{Path as AxumPath, Query, State}; -use axum::http::StatusCode; -use axum::response::IntoResponse; +use axum::http::{header::ORIGIN, HeaderMap, Request, StatusCode}; +use axum::middleware::{self, Next}; +use axum::response::{IntoResponse, Response}; use axum::routing::{get, post}; use axum::{Json, Router}; use serde::{Deserialize, Serialize}; @@ -19,11 +22,13 @@ use std::path::{Path, PathBuf}; use std::sync::Arc; use std::time::{Duration, Instant, SystemTime}; use tokio::sync::{mpsc, oneshot, Mutex, Notify}; -use tower_http::cors::CorsLayer; use uuid::Uuid; const HOST: &str = "127.0.0.1"; const API_PORT: u16 = 18767; +const LOOPBACK_API_ORIGIN: &str = "http://127.0.0.1:18767"; +const LOCALHOST_API_ORIGIN: &str = "http://localhost:18767"; +const CHROME_EXTENSION_ORIGIN_PREFIX: &str = "chrome-extension://"; // daemon 在无 CLI/API 业务请求后自动退出,避免浏览器扩展长期保持“已连接”浮层。 const IDLE_SHUTDOWN_TTL: Duration = Duration::from_secs(300); const IDLE_SHUTDOWN_CHECK_INTERVAL: Duration = Duration::from_secs(5); @@ -36,6 +41,7 @@ pub struct AppState { shutdown: mpsc::UnboundedSender<()>, sessions_ready: Arc, extension_port: u16, + api_token: String, } #[derive(Debug, Deserialize)] @@ -67,7 +73,9 @@ struct ExecRequest { #[derive(Debug, Deserialize)] struct OpenRequest { - url: String, + target: Option, + url: Option, + cwd: Option, #[serde(default = "default_active")] active: bool, switch_tab_id: Option, @@ -79,6 +87,8 @@ struct OpenRequest { window: bool, #[serde(default)] allow_focus: bool, + #[serde(default = "default_open_readiness_timeout")] + readiness_timeout: f64, } #[derive(Debug, Deserialize)] @@ -281,6 +291,10 @@ fn default_wait_timeout() -> f64 { 3.0 } +fn default_open_readiness_timeout() -> f64 { + 10.0 +} + fn default_wait_interval() -> f64 { 0.1 } @@ -332,6 +346,10 @@ impl SessionSelector { pub async fn run_daemon() -> Result<()> { let extension_port = config::load_or_create()?.extension_port; + let addr: SocketAddr = format!("{HOST}:{API_PORT}").parse()?; + let listener = tokio::net::TcpListener::bind(addr).await?; + let api_token = Uuid::new_v4().simple().to_string(); + config::save_api_token(&api_token)?; let (shutdown_tx, mut shutdown_rx) = mpsc::unbounded_channel::<()>(); let state = AppState { driver: Arc::new(Mutex::new(DriverState::default())), @@ -340,6 +358,7 @@ pub async fn run_daemon() -> Result<()> { shutdown: shutdown_tx, sessions_ready: Arc::new(Notify::new()), extension_port, + api_token, }; let ws_state = state.clone(); @@ -354,7 +373,24 @@ pub async fn run_daemon() -> Result<()> { monitor_idle_shutdown(idle_state).await; }); - let app = Router::new() + let app = build_api_router(state.clone()); + + println!("agent-browser-cli rust server listening on http://{addr}"); + let cleanup_state = state.clone(); + let serve_result = axum::serve(listener, app) + .with_graceful_shutdown(async move { + let _ = shutdown_rx.recv().await; + cleanup_on_shutdown(&cleanup_state).await; + }) + .await; + // Keep the last token file after shutdown. The next daemon rotates it after binding the API + // port, avoiding an exiting daemon racing with and deleting a newer daemon's token. + serve_result?; + Ok(()) +} + +fn build_api_router(state: AppState) -> Router { + Router::new() .route("/", get(root)) .route("/health", get(health)) .route("/tabs", get(tabs)) @@ -384,25 +420,66 @@ pub async fn run_daemon() -> Result<()> { .route("/console/clear", post(console_clear)) .route("/console/stop", post(console_stop)) .route("/shutdown", post(shutdown)) - .with_state(state.clone()) - .layer(CorsLayer::permissive()); + .layer(middleware::from_fn_with_state( + state.clone(), + require_api_access, + )) + .with_state(state) +} - let addr: SocketAddr = format!("{HOST}:{API_PORT}").parse()?; - let listener = tokio::net::TcpListener::bind(addr).await?; - println!("agent-browser-cli rust server listening on http://{addr}"); - let cleanup_state = state.clone(); - axum::serve(listener, app) - .with_graceful_shutdown(async move { - let _ = shutdown_rx.recv().await; - cleanup_on_shutdown(&cleanup_state).await; - }) - .await?; +async fn require_api_access( + State(state): State, + request: Request, + next: Next, +) -> Response { + if let Err((status, code, message)) = validate_api_access(request.headers(), &state.api_token) { + return ( + status, + Json(json!({ "ok": false, "error": message, "error_code": code })), + ) + .into_response(); + } + next.run(request).await +} + +fn validate_api_access( + headers: &HeaderMap, + expected_token: &str, +) -> std::result::Result<(), (StatusCode, &'static str, &'static str)> { + if let Some(origin) = headers.get(ORIGIN) { + let allowed = origin + .to_str() + .map(|value| matches!(value, LOOPBACK_API_ORIGIN | LOCALHOST_API_ORIGIN)) + .unwrap_or(false); + if !allowed { + return Err(( + StatusCode::FORBIDDEN, + "origin_not_allowed", + "Browser origins are not allowed to access the local API", + )); + } + } + let authorized = headers + .get(config::API_TOKEN_HEADER) + .and_then(|value| value.to_str().ok()) + .is_some_and(|value| value == expected_token); + if !authorized { + return Err(( + StatusCode::UNAUTHORIZED, + "invalid_api_token", + "A valid API token is required", + )); + } Ok(()) } +fn build_ws_router(state: AppState) -> Router { + Router::new().route("/", get(ws_handler)).with_state(state) +} + async fn run_ws_server(state: AppState) -> Result<()> { let extension_port = state.extension_port; - let app = Router::new().route("/", get(ws_handler)).with_state(state); + let app = build_ws_router(state); let addr: SocketAddr = format!("{HOST}:{extension_port}").parse()?; let listener = tokio::net::TcpListener::bind(addr).await?; println!("WebSocket server running on ws://{addr}"); @@ -410,8 +487,30 @@ async fn run_ws_server(state: AppState) -> Result<()> { Ok(()) } -async fn ws_handler(ws: WebSocketUpgrade, State(state): State) -> impl IntoResponse { +fn is_allowed_extension_origin(headers: &HeaderMap) -> bool { + let Some(origin) = headers.get(ORIGIN).and_then(|value| value.to_str().ok()) else { + return false; + }; + let Some(extension_id) = origin.strip_prefix(CHROME_EXTENSION_ORIGIN_PREFIX) else { + return false; + }; + extension_id.len() == 32 && extension_id.bytes().all(|byte| matches!(byte, b'a'..=b'p')) +} + +async fn ws_handler( + State(state): State, + headers: HeaderMap, + ws: WebSocketUpgrade, +) -> Response { + if !is_allowed_extension_origin(&headers) { + return ( + StatusCode::FORBIDDEN, + "WebSocket connections require a Chrome extension Origin", + ) + .into_response(); + } ws.on_upgrade(move |socket| handle_socket(socket, state)) + .into_response() } async fn monitor_idle_shutdown(state: AppState) { @@ -461,6 +560,13 @@ async fn handle_socket(mut socket: WebSocket, state: AppState) { } } +fn should_set_default_session(driver: &DriverState, session_key: &str) -> bool { + match driver.preferred_default_session_key.as_deref() { + Some(preferred) => preferred == session_key, + None => driver.default_session_key.is_none(), + } +} + async fn handle_ws_message( incoming: WsIncoming, state: &AppState, @@ -473,12 +579,16 @@ async fn handle_ws_message( browser_id, profile_id, profile_label, + extension_version, + file_scheme_access, tabs, } | WsIncoming::TabsUpdate { browser_id, profile_id, profile_label, + extension_version, + file_scheme_access, tabs, } => { let current: std::collections::HashSet = tabs @@ -505,8 +615,11 @@ async fn handle_ws_message( registered_ids.push(session_key.clone()); } driver.latest_session_key = Some(session_key.clone()); - if driver.default_session_key.is_none() { + if should_set_default_session(&driver, &session_key) { driver.default_session_key = Some(session_key.clone()); + if driver.preferred_default_session_key.as_deref() == Some(&session_key) { + driver.preferred_default_session_key = None; + } } driver.sessions.insert( session_key.clone(), @@ -516,6 +629,8 @@ async fn handle_ws_message( browser_id: browser_id.clone(), profile_id: profile_id.clone(), profile_label: profile_label.clone(), + extension_version: extension_version.clone(), + file_scheme_access, info, sender: sender.clone(), disconnected_at: None, @@ -572,6 +687,7 @@ async fn root() -> &'static str { async fn health(State(state): State) -> Json { let active_tabs_count = active_tabs(&state, false).await.len(); let extension_connected = has_extension_connection(&state).await; + let profile_capabilities = extension_profile_capabilities(&state).await; let ready = extension_connected && active_tabs_count > 0; let uptime = state .started_at @@ -597,7 +713,8 @@ async fn health(State(state): State) -> Json { }, "connection": { "extension_connected": extension_connected, - "active_tabs": active_tabs_count + "active_tabs": active_tabs_count, + "profiles": profile_capabilities }, "uptime": uptime, "idle_for": idle_for, @@ -961,29 +1078,210 @@ async fn exec(State(state): State, Json(req): Json) -> Js async fn open_tab(State(state): State, Json(req): Json) -> Json { touch(&state).await; + let target = match resolve_request_target(req.target.as_deref(), req.url.as_deref()) { + Ok(target) => target, + Err(err) => return Json(open_target_error_json(err)), + }; + let resolved = match resolve_open_target(target, req.cwd.as_deref()) { + Ok(resolved) => resolved, + Err(err) => return Json(open_target_error_json(err)), + }; + let mut selector = SessionSelector::new(req.switch_tab_id, req.browser, req.profile); + + if resolved.kind == TargetKind::LocalResource + && (!req.readiness_timeout.is_finite() || req.readiness_timeout <= 0.0) + { + return Json(json!({ + "ok": false, + "error": "File readiness timeout must be a finite number greater than zero", + "error_code": "invalid_readiness_timeout", + "details": { "readiness_timeout": req.readiness_timeout } + })); + } + + if resolved.kind == TargetKind::LocalResource { + let status_payload = json!({ "cmd": "status" }).to_string(); + let status_result = + match execute_page_js(&state, &status_payload, selector.clone(), true).await { + Ok(value) => value, + Err(err) => { + return Json(json!({ + "ok": false, + "error": format!("Could not query extension file capability: {err}"), + "error_code": "extension_upgrade_required", + "details": { "minimum_extension_version": "2.1" } + })) + } + }; + if let Some(executor_session_key) = status_result.get("session_key").and_then(Value::as_str) + { + selector = SessionSelector::new(Some(executor_session_key.to_string()), None, None); + } + let status = status_result + .get("js_return") + .cloned() + .unwrap_or(Value::Null); + let installed_version = status + .get("extensionVersion") + .and_then(Value::as_str) + .unwrap_or("0"); + if !version_at_least(installed_version, "2.1") { + return Json(json!({ + "ok": false, + "error": format!("Agent Browser CLI Bridge {installed_version} does not support local files; version 2.1 or newer is required"), + "error_code": "extension_upgrade_required", + "details": { + "installed_extension_version": installed_version, + "minimum_extension_version": "2.1" + } + })); + } + if status.get("fileSchemeAccess").and_then(Value::as_bool) != Some(true) { + return Json(json!({ + "ok": false, + "error": "Agent Browser CLI Bridge is not allowed to access file URLs", + "error_code": "file_scheme_access_denied", + "details": { + "extension_id": status.get("extensionId").cloned().unwrap_or(Value::Null), + "instructions": [ + "Open chrome://extensions", + "Select Agent Browser CLI Bridge", + "Enable Allow access to file URLs" + ] + } + })); + } + } + let group_title = req.group_title.or(req.session); let payload = json!({ "cmd": "openTab", - "url": normalize_url(&req.url), + "url": resolved.navigation_url.clone(), "active": req.active, "window": req.window, "allowFocus": req.allow_focus, "groupTitle": group_title, }) .to_string(); - let result = execute_page_js( - &state, - &payload, - SessionSelector::new(req.switch_tab_id, req.browser, req.profile), - true, - ) - .await; - Json(match result { - Ok(value) => json!({ "ok": true, "result": normalize_open_result(&state, value).await }), - Err(err) => json!({ "ok": false, "error": err.to_string() }), + let value = match execute_page_js(&state, &payload, selector, true).await { + Ok(value) => value, + Err(err) => return Json(json!({ "ok": false, "error": err.to_string() })), + }; + let result = normalize_open_result(&state, value.clone()).await; + let opened_tab_id = result + .get("opened_tab_id") + .and_then(Value::as_str) + .map(str::to_string); + let expected_session_key = result + .get("opened_session_key") + .and_then(Value::as_str) + .map(str::to_string); + + if resolved.kind == TargetKind::LocalResource { + let ready = if let Some(key) = expected_session_key.as_deref() { + wait_for_session_key(&state, key, req.readiness_timeout).await + } else { + false + }; + if !ready { + return Json(json!({ + "ok": false, + "error": format!("File tab was created but did not become controllable within {} seconds", req.readiness_timeout), + "error_code": "file_tab_readiness_timeout", + "opened_tab_id": opened_tab_id, + "opened_session_key": expected_session_key, + "details": { + "timeout_seconds": req.readiness_timeout, + "tab_preserved": true, + "navigation_url": resolved.navigation_url + } + })); + } + } + + if req.active { + let mut driver = state.driver.lock().await; + if let Some(key) = expected_session_key.clone() { + if driver.sessions.get(&key).is_some_and(Session::is_active) { + driver.default_session_key = Some(key); + driver.preferred_default_session_key = None; + } else { + driver.preferred_default_session_key = Some(key); + } + } + } + Json(json!({ "ok": true, "result": result })) +} + +fn resolve_request_target<'a>( + target: Option<&'a str>, + url: Option<&'a str>, +) -> Result<&'a str, OpenTargetError> { + match (target, url) { + (Some(target), Some(url)) if target != url => Err(OpenTargetError { + code: "ambiguous_open_target", + message: "Open request contains conflicting target and url fields".to_string(), + details: json!({ "target": target, "url": url }), + }), + (Some(target), _) => Ok(target), + (_, Some(url)) => Ok(url), + (None, None) => Err(OpenTargetError { + code: "missing_open_target", + message: "Open request requires target or url".to_string(), + details: json!({}), + }), + } +} + +fn open_target_error_json(err: OpenTargetError) -> Value { + json!({ + "ok": false, + "error": err.message, + "error_code": err.code, + "details": err.details, }) } +fn version_at_least(installed: &str, minimum: &str) -> bool { + let parse = |value: &str| { + value + .split('.') + .map(|part| part.parse::().unwrap_or(0)) + .collect::>() + }; + let mut installed = parse(installed); + let mut minimum = parse(minimum); + let len = installed.len().max(minimum.len()); + installed.resize(len, 0); + minimum.resize(len, 0); + installed >= minimum +} + +async fn wait_for_session_key(state: &AppState, session_key: &str, timeout: f64) -> bool { + let deadline = Instant::now() + Duration::from_secs_f64(timeout.max(0.1)); + loop { + { + let driver = state.driver.lock().await; + if driver + .sessions + .get(session_key) + .is_some_and(Session::is_active) + { + return true; + } + } + let now = Instant::now(); + if now >= deadline { + return false; + } + let remaining = deadline.saturating_duration_since(now); + tokio::select! { + _ = state.sessions_ready.notified() => {} + _ = tokio::time::sleep(remaining.min(Duration::from_millis(100))) => {} + } + } +} + async fn normalize_open_result(state: &AppState, value: Value) -> Value { let opened = value.get("js_return").cloned().unwrap_or(Value::Null); let opened_tab_id = opened.get("id").and_then(|v| { @@ -996,11 +1294,10 @@ async fn normalize_open_result(state: &AppState, value: Value) -> Value { .and_then(Value::as_str) .map(str::to_string); let opened_session_key = if let Some(tab_id) = opened_tab_id.as_deref() { - find_session_key_by_tab_id(state, tab_id).await.or_else(|| { - executor_session_key - .as_deref() - .and_then(|key| derive_session_key(key, tab_id)) - }) + match executor_session_key.as_deref() { + Some(executor_key) => derive_session_key(executor_key, tab_id), + None => find_unique_session_key_by_tab_id(state, tab_id).await, + } } else { None }; @@ -1024,19 +1321,28 @@ async fn normalize_open_result(state: &AppState, value: Value) -> Value { }) } -async fn find_session_key_by_tab_id(state: &AppState, tab_id: &str) -> Option { +async fn find_unique_session_key_by_tab_id(state: &AppState, tab_id: &str) -> Option { let driver = state.driver.lock().await; - driver + let mut matches = driver .sessions .values() - .find(|session| session.is_active() && session.tab_id == tab_id) - .map(|session| session.session_key.clone()) + .filter(|session| session.is_active() && session.tab_id == tab_id) + .map(|session| session.session_key.clone()); + let session_key = matches.next()?; + if matches.next().is_some() { + return None; + } + Some(session_key) } fn derive_session_key(executor_session_key: &str, opened_tab_id: &str) -> Option { let mut parts = executor_session_key.splitn(3, ':'); let browser_id = parts.next()?; let profile_id = parts.next()?; + let executor_tab_id = parts.next()?; + if browser_id.is_empty() || profile_id.is_empty() || executor_tab_id.is_empty() { + return None; + } Some(crate::protocol::make_session_key( browser_id, profile_id, @@ -1306,6 +1612,36 @@ async fn active_tabs_filtered( .collect()) } +async fn extension_profile_capabilities(state: &AppState) -> Vec { + let driver = state.driver.lock().await; + let mut profiles: HashMap<(String, String), Value> = HashMap::new(); + for session in driver + .sessions + .values() + .filter(|session| session.is_active()) + { + profiles + .entry((session.browser_id.clone(), session.profile_id.clone())) + .or_insert_with(|| { + json!({ + "browser_id": session.browser_id, + "profile_id": session.profile_id, + "profile_label": session.profile_label, + "extension_version": session.extension_version, + "file_scheme_access": session.file_scheme_access, + }) + }); + } + let mut values: Vec = profiles.into_values().collect(); + values.sort_by(|a, b| { + a.get("profile_id") + .and_then(Value::as_str) + .unwrap_or("") + .cmp(b.get("profile_id").and_then(Value::as_str).unwrap_or("")) + }); + values +} + async fn has_extension_connection(state: &AppState) -> bool { let driver = state.driver.lock().await; driver @@ -1383,10 +1719,7 @@ async fn execute_page_js( selector: SessionSelector, no_monitor: bool, ) -> Result { - if selector.tab_id.is_some() || selector.browser.is_some() || selector.profile.is_some() { - let session_key = select_tab(state, selector.clone()).await?; - state.driver.lock().await.default_session_key = Some(session_key); - } + let executor_session_key = select_tab(state, selector).await?; let before = if no_monitor { None } else { @@ -1397,13 +1730,16 @@ async fn execute_page_js( .into_iter() .map(|t| t.session_key) .collect(); - let response = execute_raw_js(state, script, Duration::from_secs(15)).await?; - let current_session_key = state.driver.lock().await.default_session_key.clone(); - let current_tab_id = if let Some(key) = current_session_key.as_deref() { - session_tab_id(state, key).await.ok() - } else { - None - }; + let execution = execute_raw_js_with_session( + state, + script, + Duration::from_secs(15), + Some(&executor_session_key), + ) + .await?; + let current_session_key = execution.session_key; + let current_tab_id = execution.tab_id; + let response = execution.result; let mut result = json!({ "status": "success", "js_return": response.data.or(response.result).unwrap_or(Value::Null), @@ -1433,37 +1769,63 @@ async fn execute_page_js( Ok(result) } +struct ExecWithSession { + result: ExecResult, + session_key: String, + tab_id: String, +} + async fn execute_raw_js(state: &AppState, code: &str, timeout: Duration) -> Result { + Ok(execute_raw_js_with_session(state, code, timeout, None) + .await? + .result) +} + +async fn execute_raw_js_with_session( + state: &AppState, + code: &str, + timeout: Duration, + requested_session_key: Option<&str>, +) -> Result { let (session_key, tab_id, sender) = { wait_for_sessions(state, Duration::from_secs(5)).await; let driver = state.driver.lock().await; - let session_key = driver - .default_session_key - .as_ref() - .and_then(|key| { - driver - .sessions - .get(key) - .filter(|s| s.is_active()) - .map(|s| s.session_key.clone()) - }) - .or_else(|| { - driver.latest_session_key.as_ref().and_then(|key| { + let session_key = if let Some(requested) = requested_session_key { + driver + .sessions + .get(requested) + .filter(|session| session.is_active()) + .map(|session| session.session_key.clone()) + .ok_or_else(|| anyhow!("会话ID {requested} 未连接"))? + } else { + driver + .default_session_key + .as_ref() + .and_then(|key| { driver .sessions .get(key) .filter(|s| s.is_active()) .map(|s| s.session_key.clone()) }) - }) - .or_else(|| { - driver - .sessions - .values() - .find(|s| s.is_active()) - .map(|s| s.session_key.clone()) - }) - .ok_or_else(|| anyhow!("没有可用的浏览器标签页,查L3记忆分析原因。"))?; + .or_else(|| { + driver.latest_session_key.as_ref().and_then(|key| { + driver + .sessions + .get(key) + .filter(|s| s.is_active()) + .map(|s| s.session_key.clone()) + }) + }) + .or_else(|| { + driver + .sessions + .values() + .find(|s| s.is_active()) + .map(|s| s.session_key.clone()) + }) + .ok_or_else(|| anyhow!("没有可用的浏览器标签页,查L3记忆分析原因。"))? + }; let session = driver .sessions .get(&session_key) @@ -1492,7 +1854,7 @@ async fn execute_raw_js(state: &AppState, code: &str, timeout: Duration) -> Resu sender .send(payload) .map_err(|_| anyhow!("浏览器扩展连接已断开"))?; - match tokio::time::timeout(timeout, rx).await { + let result = match tokio::time::timeout(timeout, rx).await { Ok(Ok(value)) => value, Ok(Err(_)) => Err(anyhow!("执行结果通道已关闭")), Err(_) => { @@ -1522,7 +1884,12 @@ async fn execute_raw_js(state: &AppState, code: &str, timeout: Duration) -> Resu }) } } - } + }?; + Ok(ExecWithSession { + result, + session_key, + tab_id, + }) } async fn get_html( @@ -1625,14 +1992,6 @@ return {{ result: __mainResult, wait: {{ ok: __matched, matched: __matched, valu ) } -fn normalize_url(url: &str) -> String { - if url.starts_with("http://") || url.starts_with("https://") { - url.to_string() - } else { - format!("https://{url}") - } -} - async fn touch(state: &AppState) { *state.last_activity.lock().await = Instant::now(); } @@ -2928,11 +3287,12 @@ fn parse_key_segment(segment: &str, platform: &str) -> Result { modifiers.push(spec); } let mut key = key_spec(parts[parts.len() - 1])?; - if modifier_bits & 8 != 0 { - if key.key.len() == 1 && key.key.chars().all(|c| c.is_ascii_lowercase()) { - key.key = key.key.to_uppercase(); - key.text = key.text.as_ref().map(|_| key.key.clone()); - } + if modifier_bits & 8 != 0 + && key.key.len() == 1 + && key.key.chars().all(|c| c.is_ascii_lowercase()) + { + key.key = key.key.to_uppercase(); + key.text = key.text.as_ref().map(|_| key.key.clone()); } Ok(KeySegment { modifiers, @@ -3117,13 +3477,7 @@ fn ax_string(value: &Option) -> Option { } fn non_empty(value: Option) -> Option { - value.and_then(|value| { - if value.trim().is_empty() { - None - } else { - Some(value) - } - }) + value.filter(|value| !value.trim().is_empty()) } fn truncate_text(value: &str, max: usize) -> String { @@ -3661,3 +4015,358 @@ fn decode_base64(input: &str) -> Result> { } Ok(out) } + +#[cfg(test)] +mod tests { + use super::*; + use axum::http::HeaderValue; + + fn test_state() -> AppState { + let (shutdown, _shutdown_rx) = mpsc::unbounded_channel(); + AppState { + driver: Arc::new(Mutex::new(DriverState::default())), + started_at: SystemTime::now(), + last_activity: Arc::new(Mutex::new(Instant::now())), + shutdown, + sessions_ready: Arc::new(Notify::new()), + extension_port: config::DEFAULT_EXTENSION_PORT, + api_token: "test-token".to_string(), + } + } + + fn insert_test_session( + driver: &mut DriverState, + browser_id: &str, + profile_id: &str, + tab_id: &str, + ) -> mpsc::UnboundedReceiver { + let session_key = crate::protocol::make_session_key(browser_id, profile_id, tab_id); + let (sender, receiver) = mpsc::unbounded_channel(); + driver.sessions.insert( + session_key.clone(), + Session { + session_key: session_key.clone(), + tab_id: tab_id.to_string(), + browser_id: browser_id.to_string(), + profile_id: profile_id.to_string(), + profile_label: None, + extension_version: Some("2.1".to_string()), + file_scheme_access: Some(true), + info: TabInfo { + id: tab_id.to_string(), + tab_id: tab_id.to_string(), + browser_id: browser_id.to_string(), + profile_id: profile_id.to_string(), + profile_label: None, + session_key, + url: "https://example.com".to_string(), + title: "test".to_string(), + tab_type: "ext_ws".to_string(), + connected_at: None, + }, + sender, + disconnected_at: None, + }, + ); + receiver + } + + #[test] + fn api_access_requires_token_and_rejects_foreign_origins() { + let mut headers = HeaderMap::new(); + headers.insert( + config::API_TOKEN_HEADER, + HeaderValue::from_static("test-token"), + ); + assert!(validate_api_access(&headers, "test-token").is_ok()); + + headers.insert(ORIGIN, HeaderValue::from_static("https://evil.example")); + let error = validate_api_access(&headers, "test-token").unwrap_err(); + assert_eq!(error.0, StatusCode::FORBIDDEN); + assert_eq!(error.1, "origin_not_allowed"); + + headers.insert(ORIGIN, HeaderValue::from_static(LOOPBACK_API_ORIGIN)); + assert!(validate_api_access(&headers, "test-token").is_ok()); + headers.remove(config::API_TOKEN_HEADER); + let error = validate_api_access(&headers, "test-token").unwrap_err(); + assert_eq!(error.0, StatusCode::UNAUTHORIZED); + assert_eq!(error.1, "invalid_api_token"); + } + + #[test] + fn websocket_origin_requires_a_valid_chrome_extension_id() { + let mut headers = HeaderMap::new(); + assert!(!is_allowed_extension_origin(&headers)); + + headers.insert(ORIGIN, HeaderValue::from_static("https://evil.example")); + assert!(!is_allowed_extension_origin(&headers)); + + headers.insert( + ORIGIN, + HeaderValue::from_static("chrome-extension://abcdefghijklmnopabcdefghijklmnop"), + ); + assert!(is_allowed_extension_origin(&headers)); + + headers.insert( + ORIGIN, + HeaderValue::from_static( + "chrome-extension://abcdefghijklmnopabcdefghijklmnop.example.com", + ), + ); + assert!(!is_allowed_extension_origin(&headers)); + } + + #[tokio::test(flavor = "multi_thread", worker_threads = 2)] + async fn websocket_router_rejects_missing_and_non_extension_origins() { + let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap(); + let addr = listener.local_addr().unwrap(); + let server = tokio::spawn(async move { + axum::serve(listener, build_ws_router(test_state())) + .await + .unwrap(); + }); + + let status_lines = tokio::task::spawn_blocking(move || { + use std::io::{BufRead, BufReader, Write}; + + let handshake = |origin: Option<&str>| { + let mut stream = std::net::TcpStream::connect(addr).unwrap(); + stream + .set_read_timeout(Some(Duration::from_secs(2))) + .unwrap(); + let origin_header = origin + .map(|value| format!("Origin: {value}\r\n")) + .unwrap_or_default(); + let request = format!( + "GET / HTTP/1.1\r\n\ + Host: {addr}\r\n\ + Upgrade: websocket\r\n\ + Connection: Upgrade\r\n\ + Sec-WebSocket-Key: dGhlIHNhbXBsZSBub25jZQ==\r\n\ + Sec-WebSocket-Version: 13\r\n\ + {origin_header}\r\n" + ); + stream.write_all(request.as_bytes()).unwrap(); + stream.flush().unwrap(); + let mut status_line = String::new(); + BufReader::new(stream).read_line(&mut status_line).unwrap(); + status_line + }; + + ( + handshake(None), + handshake(Some("https://evil.example")), + handshake(Some("chrome-extension://abcdefghijklmnopabcdefghijklmnop")), + ) + }) + .await + .unwrap(); + server.abort(); + let _ = server.await; + + assert!(status_lines.0.starts_with("HTTP/1.1 403 Forbidden")); + assert!(status_lines.1.starts_with("HTTP/1.1 403 Forbidden")); + assert!(status_lines + .2 + .starts_with("HTTP/1.1 101 Switching Protocols")); + } + + #[tokio::test(flavor = "multi_thread", worker_threads = 2)] + async fn api_router_enforces_token_and_origin_checks() { + let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap(); + let addr = listener.local_addr().unwrap(); + let server = tokio::spawn(async move { + axum::serve(listener, build_api_router(test_state())) + .await + .unwrap(); + }); + + let responses = tokio::task::spawn_blocking(move || { + let client = reqwest::blocking::Client::new(); + let url = format!("http://{addr}/health"); + let send = |token: Option<&str>, origin: Option<&str>| { + let mut request = client.get(&url); + if let Some(token) = token { + request = request.header(config::API_TOKEN_HEADER, token); + } + if let Some(origin) = origin { + request = request.header("origin", origin); + } + let response = request.send().unwrap(); + let status = response.status(); + let body = response.json::().unwrap(); + (status, body) + }; + ( + send(None, None), + send(Some("wrong-token"), None), + send(Some("test-token"), Some("https://evil.example")), + send(Some("test-token"), None), + ) + }) + .await + .unwrap(); + server.abort(); + let _ = server.await; + + assert_eq!(responses.0 .0, StatusCode::UNAUTHORIZED); + assert_eq!(responses.0 .1["error_code"], "invalid_api_token"); + assert_eq!(responses.1 .0, StatusCode::UNAUTHORIZED); + assert_eq!(responses.1 .1["error_code"], "invalid_api_token"); + assert_eq!(responses.2 .0, StatusCode::FORBIDDEN); + assert_eq!(responses.2 .1["error_code"], "origin_not_allowed"); + assert_eq!(responses.3 .0, StatusCode::OK); + assert_eq!(responses.3 .1["running"], true); + } + + #[tokio::test] + async fn open_result_uses_executor_scope_when_browser_tab_ids_overlap() { + let state = test_state(); + { + let mut driver = state.driver.lock().await; + let _receiver_a = insert_test_session(&mut driver, "browser-a", "profile-a", "7"); + let _receiver_b = insert_test_session(&mut driver, "browser-b", "profile-b", "42"); + } + + let result = normalize_open_result( + &state, + json!({ + "js_return": { "id": 42 }, + "tab_id": "7", + "session_key": "browser-a:profile-a:7" + }), + ) + .await; + assert_eq!( + result.get("opened_session_key").and_then(Value::as_str), + Some("browser-a:profile-a:42") + ); + } + + #[tokio::test] + async fn open_result_refuses_ambiguous_tab_id_without_executor_scope() { + let state = test_state(); + { + let mut driver = state.driver.lock().await; + let _receiver_a = insert_test_session(&mut driver, "browser-a", "profile-a", "42"); + let _receiver_b = insert_test_session(&mut driver, "browser-b", "profile-b", "42"); + } + + let result = normalize_open_result(&state, json!({ "js_return": { "id": 42 } })).await; + assert!(result.get("opened_session_key").is_some_and(Value::is_null)); + } + + #[tokio::test] + async fn page_execution_reports_the_session_that_received_the_command() { + let state = test_state(); + let (mut receiver_a, _receiver_b) = { + let mut driver = state.driver.lock().await; + let receiver_a = insert_test_session(&mut driver, "browser-a", "profile-a", "7"); + let receiver_b = insert_test_session(&mut driver, "browser-b", "profile-b", "8"); + driver.default_session_key = Some("browser-a:profile-a:7".to_string()); + (receiver_a, receiver_b) + }; + + let exec_state = state.clone(); + let execution = tokio::spawn(async move { + execute_page_js( + &exec_state, + "return 1", + SessionSelector::new(Some("browser-a:profile-a:7".to_string()), None, None), + true, + ) + .await + }); + let payload = tokio::time::timeout(Duration::from_secs(1), receiver_a.recv()) + .await + .expect("command should be dispatched") + .expect("browser-a receiver should remain connected"); + let exec_id = serde_json::from_str::(&payload) + .unwrap() + .get("id") + .and_then(Value::as_str) + .unwrap() + .to_string(); + + state.driver.lock().await.default_session_key = Some("browser-b:profile-b:8".to_string()); + let (sender, _receiver) = mpsc::unbounded_channel(); + handle_ws_message( + WsIncoming::Result { + id: exec_id, + result: json!(1), + new_tabs: None, + }, + &state, + sender, + &mut Vec::new(), + ) + .await; + + let result = execution.await.unwrap().unwrap(); + assert_eq!( + result.get("session_key").and_then(Value::as_str), + Some("browser-a:profile-a:7") + ); + assert_eq!(result.get("tab_id").and_then(Value::as_str), Some("7")); + } + + #[test] + fn pending_default_session_does_not_match_another_browser_tab() { + let mut driver = DriverState { + preferred_default_session_key: Some("browser-a:profile-a:42".to_string()), + ..DriverState::default() + }; + assert!(!should_set_default_session( + &driver, + "browser-b:profile-b:42" + )); + assert!(should_set_default_session( + &driver, + "browser-a:profile-a:42" + )); + driver.default_session_key = Some("browser-a:profile-a:7".to_string()); + assert!(should_set_default_session( + &driver, + "browser-a:profile-a:42" + )); + } + + #[test] + fn request_target_accepts_compatible_fields() { + assert_eq!( + resolve_request_target(Some("./demo.html"), None).unwrap(), + "./demo.html" + ); + assert_eq!( + resolve_request_target(None, Some("https://example.com")).unwrap(), + "https://example.com" + ); + assert_eq!( + resolve_request_target(Some("same"), Some("same")).unwrap(), + "same" + ); + } + + #[test] + fn request_target_rejects_missing_and_conflicting_fields() { + assert_eq!( + resolve_request_target(None, None).unwrap_err().code, + "missing_open_target" + ); + assert_eq!( + resolve_request_target(Some("a"), Some("b")) + .unwrap_err() + .code, + "ambiguous_open_target" + ); + } + + #[test] + fn extension_versions_compare_numerically() { + assert!(version_at_least("2.1", "2.1")); + assert!(version_at_least("2.10", "2.1")); + assert!(version_at_least("2.1.1", "2.1")); + assert!(!version_at_least("2.0", "2.1")); + assert!(!version_at_least("unknown", "2.1")); + } +}