Compare commits

..
Author SHA1 Message Date
leeguooooo 8880aa2f35 chore(release): 0.16.3-fork.2
- restore extension mode default to headed when headless is unspecified

- keep explicit headless override behavior
2026-03-05 11:44:47 +09:00
leeguooooo f051e72f85 chore(release): bump to 0.16.3-fork.1 2026-03-05 11:20:06 +09:00
leeguooooo 6ae703565c fix(connection): surface daemon startup stderr during launch 2026-03-05 11:18:21 +09:00
layla d56442cf91 Fix dialog dismiss command parsing (#605) 2026-03-05 11:16:14 +09:00
Li Yang dad7be8c77 fix: use reqwest for CDP port discovery instead of broken hand-rolled HTTP client (#619)
reqwest_get_string() was hand-rolling HTTP/1.1 over raw TCP despite reqwest
being an existing dependency. The hand-rolled implementation had two bugs:

1. URL path parsing: url.find('/') matched the first '/' in 'http://',
   producing path '//127.0.0.1:9222/json/version' instead of '/json/version'

2. read_to_end() hangs: Chrome's DevTools HTTP server ignores Connection: close
   and keeps the socket open, so read_to_end() waits for EOF that never comes

This caused 'agent-browser --cdp <port>' to always timeout when AGENT_BROWSER_NATIVE=1.

Fix: replace 49 lines of broken TCP code with reqwest::get(), which was
already in Cargo.toml.
2026-03-05 11:16:14 +09:00
Chris Tate 7921928ec4 headed mode (#607)
* headed mode

* fixes

* fixes

* docs

* fixes

* fixes

* fixes
2026-03-05 11:15:17 +09:00
leeguooooo a71198d591 fix(stealth): neutralize disable-devtool auto bootstrap 2026-03-05 11:11:36 +09:00
leeguooooo 19aefb7eb8 Update binaries for 0.16.1-fork.4 2026-03-04 18:46:17 +09:00
leeguooooo ff2e66c0a1 chore(release): bump to 0.16.1-fork.4 2026-03-04 18:24:10 +09:00
leeguooooo d320e1df47 refactor(cli): ignore --session and enforce default runtime session 2026-03-04 18:19:33 +09:00
leeguooooo 7b454e14ed fix: stabilize native release builds across platforms 2026-03-04 15:44:28 +09:00
leeguooooo 17d44baa2f fix(release): enforce binary version integrity pre/post publish 2026-03-04 15:01:52 +09:00
leeguooooo 59589043db fix(stealth): bypass worker wrapping in cloudflare challenge runtime 2026-03-04 13:48:16 +09:00
leeguooooo 9e7c8937a6 chore: 提交剩余文档与扩展面板资源更新 2026-03-04 11:54:55 +09:00
leeguooooo 2d0d18f0d9 fix(extension): default to high-risk mode by disabling page bridge 2026-03-04 11:46:44 +09:00
leeguooooo ae22bda46c chore(release): bump to 0.16.1-fork.3 2026-03-04 11:35:11 +09:00
leeguooooo d82357fea4 fix(release): enforce bundled binary version checks and bump to 0.16.1-fork.2 2026-03-04 10:25:04 +09:00
leeguooooo 66a39f3c83 feat: add abs alias, refresh extension panel, and bump to 0.16.1-fork.1 2026-03-04 09:57:56 +09:00
leeguooooo eedf824af7 merge: upstream v0.16.1 while preserving fork stealth behaviors 2026-03-04 09:48:58 +09:00
leeguooooo c07eb7ee52 chore: 更新 Cloudflare 及浏览器自动化攻防文章并补发 blog 链接 2026-03-04 09:39:38 +09:00
github-actions[bot]andgithub-actions[bot] <41898282+github-actions[bot]@users.noreply.github.com> a493d02c66 chore: version packages (#604)
Co-authored-by: github-actions[bot] <41898282+github-actions[bot]@users.noreply.github.com>
2026-03-03 17:59:02 -06:00
Chris Tate c4180c8cb1 chore: add patch changeset for release (#603) 2026-03-03 17:51:29 -06:00
Chris Tate 56260f68b0 Native: auto-detect sandbox/container environments for Chrome launch (#602)
Fixes #600

Three improvements to `--native` Chrome launching:

- `find_chrome()` now falls back to Playwright's browser cache (`~/.cache/ms-playwright/`) when no system Chrome is found
- Auto-detect containers/VMs (root, Docker, Podman, cgroups) and inject `--no-sandbox`
- Chrome stderr is now captured and included in launch error messages, with a hint when sandbox errors are detected
2026-03-03 17:45:23 -06:00
Chris Tate 324a9e4e0c windows 8 cores (#599)
* windows 8 cores

* add workflow dispatch
2026-03-03 17:26:22 -06:00
Chris Tate 7f42eed031 faster ci (#598) 2026-03-03 16:45:57 -06:00
github-actions[bot]andgithub-actions[bot] <41898282+github-actions[bot]@users.noreply.github.com> c10981413f chore: version packages (#597)
Co-authored-by: github-actions[bot] <41898282+github-actions[bot]@users.noreply.github.com>
2026-03-03 16:18:08 -06:00
Chris Tate 05018b309a prepare v0.16.0 (#596) 2026-03-03 16:09:39 -06:00
Chris Tate 9d0454d229 fix: switch from native-tls to rustls for cross-compilation (#595)
The native PR introduced tokio-tungstenite and reqwest with native-tls,
which depends on openssl-sys (C library). This breaks the release
workflow's cargo-zigbuild cross-compilation on Linux because zig's C
compiler can't find the system OpenSSL headers.

Switch to rustls (pure Rust TLS) which has zero C dependencies and
cross-compiles trivially. Also shrinks the dependency tree.
2026-03-03 15:45:20 -06:00
Chris Tate 51f5fa484c native (#594)
* Native Rust rewrite of agent-browser daemon

Single-binary Rust implementation replacing the Node.js/Playwright daemon
with direct CDP (Chrome DevTools Protocol) communication. Includes full
command parity, WebDriver/Safari/iOS backend routing, request tracking,
frame context management, CDP protocol codegen, and comprehensive tests.

* improvements

* fix ci

* fixes

* faster builds
2026-03-03 15:15:57 -06:00
github-actions[bot]andgithub-actions[bot] <41898282+github-actions[bot]@users.noreply.github.com> 857c0b25df chore: version packages (#591)
Co-authored-by: github-actions[bot] <41898282+github-actions[bot]@users.noreply.github.com>
2026-03-03 08:21:58 -06:00
Chris Tate 62241b50e9 chore: add patch changeset for release (#589) 2026-03-03 08:08:02 -06:00
Chris Tate c6a33b6338 fix(windows): resolve daemon startup failures and Git Bash compatibility (#582)
* fix(windows): resolve daemon startup failures and Git Bash compatibility

Three root causes behind 27 open Windows issues:

1. Path::canonicalize() returns \\?\ prefixed paths on Windows that
   Node.js cannot parse, preventing daemon startup. Strip the prefix
   before passing to Node. (fixes #522, #390, #56, #25, #37, #89)

2. Git Bash/MSYS2 translates Unix-style paths and resolves node to
   a shell wrapper script. Use node.exe explicitly and set
   MSYS_NO_PATHCONV/MSYS2_ARG_CONV_EXCL to prevent argument mangling.
   (fixes #148, #108, #171)

3. postinstall fixWindowsShims() hardcoded x64 arch and did not verify
   the native binary exists before rewriting shims. Now detects arch
   dynamically and validates the binary path. (fixes #262)

Also:
- Error messages now show TCP port on Windows instead of Unix socket path
- Windows CI expanded to test full daemon lifecycle (open, snapshot, close)

* fix(windows): strip \\?\ prefix in auth-cli path (fixes #579)

Same canonicalize() issue as the daemon spawn path, but in
run_auth_cli() which passes the script path to Node.js.
2026-03-03 08:00:14 -06:00
leeguooooo d32a1d046a Detail agent stealth Cloudflare fix 2026-03-03 18:04:44 +09:00
leeguooooo 0c0ed5e72c chore(release): bump version to 0.15.2-fork.2 2026-03-03 17:13:07 +09:00
leeguooooo 34092ec193 fix: prevent stale native version and skip npm-bin rewrite on pnpm 2026-03-03 17:12:17 +09:00
leeguooooo 726377c4c1 feat: add doctor diagnostics and bump to 0.15.2-fork.1 2026-03-03 17:05:42 +09:00
leeguooooo 8e2e4abce6 feat: complete agent-browser-stealth extension controls 2026-03-03 12:33:46 +09:00
leeguooooo 870895e922 feat: expand agent-browser-stealth extension capabilities 2026-03-03 12:28:34 +09:00
leeguooooo 2a766cfe48 feat: add CDP tab-group plugin handshake with silent fallback 2026-03-03 12:17:14 +09:00
leeguooooo d04cf59238 feat: auto-group agent tabs by default in local Chromium 2026-03-03 11:37:53 +09:00
leeguooooo 0a257ad2c1 merge: sync upstream/main into fork main (v0.15.2) 2026-03-03 09:57:17 +09:00
github-actions[bot]andgithub-actions[bot] <41898282+github-actions[bot]@users.noreply.github.com> d97e2016f5 chore: version packages (#585)
Co-authored-by: github-actions[bot] <41898282+github-actions[bot]@users.noreply.github.com>
2026-03-02 17:16:52 -06:00
Chris Tate 6aea316c82 chore: add patch changeset for release (#583) 2026-03-02 17:10:32 -06:00
Chris Tate c7fa10cb1b remove skill creator (#581) 2026-03-02 16:26:19 -06:00
leeguooooo 44c0361fcd Update readme with stealth FAQ 2026-03-02 18:32:28 +09:00
leeguooooo 907ca8c808 Update agent-browser skill installs 2026-03-02 09:38:48 +09:00
Giulio Leone b304a4188c fix: correct misleading output for cookies clear and tab close (#556) (#563)
Bug 1: `cookies clear` printed 'Request log cleared' instead of 'Cookies cleared'
because the output handler matched the generic `{ cleared: true }` response shape
without checking the action context. Now uses the `action` parameter to distinguish
`cookies_clear` from `requests --clear`.

Bug 2: `tab close` printed 'Browser closed' instead of 'Tab closed' because the
output handler matched the generic `{ closed: ... }` response shape without checking
the action context. Now uses the `action` parameter to distinguish `tab_close` from
`close` (full browser close).

Closes #556
2026-03-01 12:23:23 -06:00
neilmixandClaude Opus 4.6 e912f541f2 fix: treat EPERM from kill(pid, 0) as "process exists" in daemon liveness checks (#564)
Per POSIX, kill(pid, 0) returns EPERM when the process exists but the
caller lacks permission to signal it, and ESRCH when it does not exist.
The daemon liveness checks in both the Rust CLI and TypeScript daemon
treated any kill failure as "not running", which is incorrect when
running inside a macOS sandbox that restricts signal delivery to
(target self). This caused the CLI to delete the real daemon's socket
and PID files, then spawn a duplicate daemon.

Co-authored-by: Claude Opus 4.6 <noreply@anthropic.com>
2026-03-01 10:12:37 -06:00
7238b7da4c fix: resolve unnamed element refs matching multiple elements (#573)
* fix: resolve unnamed element refs matching multiple elements (#500)

When a page has one unnamed button among several named buttons,
clicking its ref fails with "matched N elements" because the
locator `getByRole('button')` matches all buttons on the page.

Normalize unnamed interactive elements to `name: ""` so the
selector becomes `getByRole('button', { name: "", exact: true })`
which matches only buttons with empty accessible names.

Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>

* refactor: remove dead code branch in buildSelector

Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>

* refactor: make RefMap.name required string, remove dead code branches

Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>

---------

Co-authored-by: hyunjinee <leehj0110@kakao.com>
Co-authored-by: Claude Opus 4.6 <noreply@anthropic.com>
2026-03-01 09:42:05 -06:00
Chris Tate 79d8dfe34c add skills to docs (#576) 2026-03-01 09:02:06 -06:00
Chris Tate 14ec5b5ffa add slack skill (#571) 2026-02-28 12:03:50 -06:00
leeguooooo 74fda70b67 fix(release): publish correct native version in fork.7 2026-02-28 11:29:21 +09:00
leeguooooo 2a397de59f merge: sync upstream/main into fork main 2026-02-28 11:19:50 +09:00
leeguooooo bf672ee7f9 fix(stealth): avoid matchMedia Illegal invocation
- bind MediaQueryList methods when proxied for prefers-color-scheme light patch

- add regression test for addEventListener/removeEventListener

- bump version to 0.14.0-fork.6 and sync Cargo metadata
2026-02-28 11:17:02 +09:00
leeguooooo 74910cfef1 Merge tag 'v0.15.0' into codex/sync-v0.15.0
v0.15.0

# Conflicts:
#	CHANGELOG.md
#	README.md
#	cli/Cargo.lock
#	cli/Cargo.toml
#	cli/src/commands.rs
#	cli/src/connection.rs
#	cli/src/flags.rs
#	cli/src/main.rs
#	docs/src/app/commands/page.mdx
#	docs/src/app/configuration/page.mdx
#	package.json
#	src/actions.ts
2026-02-27 10:49:55 +09:00
leeguooooo 41830dff71 fix(cookies): require domain and path together when url is absent 2026-02-27 10:49:12 +09:00
github-actions[bot]andgithub-actions[bot] <41898282+github-actions[bot]@users.noreply.github.com> 79b05877a8 chore: version packages (#548)
Co-authored-by: github-actions[bot] <41898282+github-actions[bot]@users.noreply.github.com>
2026-02-26 11:45:22 -06:00
Chris Tate 7bd8ce937b chore: add patch changeset for release (#546) 2026-02-26 11:36:47 -06:00
Ryan Siddle b455a58aa2 fix: preserve chrome-extension:// and chrome:// URL schemes in CLI (#410)
The CLI's URL normalization was auto-prepending https:// to any URL
whose scheme wasn't in the allowlist (http, https, about, data, file).
This caused chrome-extension:// URLs to become
https://chrome-extension//... which fails with ERR_NAME_NOT_RESOLVED,
preventing navigation to extension pages (popup, side panel, options).

Add chrome-extension:// and chrome:// to the open command's scheme
allowlist, and update the record start/restart commands to preserve
any URL that already contains :// instead of only checking for http.

Fixes #409
2026-02-26 11:30:02 -06:00
github-actions[bot]andgithub-actions[bot] <41898282+github-actions[bot]@users.noreply.github.com> b59dc4c82c chore: version packages (#545)
Co-authored-by: github-actions[bot] <41898282+github-actions[bot]@users.noreply.github.com>
2026-02-25 15:57:43 -06:00
Chris Tate 2e38882664 prepare v0.15 (#544)
* add security hardening features

- Add authentication vault (`auth save/login/list/show/delete`) so credentials are stored locally and never exposed to the LLM (fixes Snyk W007)
- Add `--content-boundaries` flag to wrap page-sourced output in structural markers, helping LLMs distinguish tool output from untrusted page content (fixes Snyk W011)
- Add `--allowed-domains` flag to restrict browser navigation to trusted domains
- Add `--action-policy` for static allow/deny gating of action categories, with opt-in `--confirm-actions`/`--confirm-interactive` for orchestrator or human-in-the-loop confirmation
- Add `--max-output` flag to truncate large page outputs, preventing context flooding
- New docs page at /security, updated README, SKILL.md, CLI help text, and templates

* fixes

* fixes

* fixes

* fixes

* fixes

* fixes

* fixes

* docs

* prepare v0.15
2026-02-25 15:47:26 -06:00
Chris Tate bc1e917e87 add security hardening features (#543)
* add security hardening features

- Add authentication vault (`auth save/login/list/show/delete`) so credentials are stored locally and never exposed to the LLM (fixes Snyk W007)
- Add `--content-boundaries` flag to wrap page-sourced output in structural markers, helping LLMs distinguish tool output from untrusted page content (fixes Snyk W011)
- Add `--allowed-domains` flag to restrict browser navigation to trusted domains
- Add `--action-policy` for static allow/deny gating of action categories, with opt-in `--confirm-actions`/`--confirm-interactive` for orchestrator or human-in-the-loop confirmation
- Add `--max-output` flag to truncate large page outputs, preventing context flooding
- New docs page at /security, updated README, SKILL.md, CLI help text, and templates

* fixes

* fixes

* fixes

* fixes

* fixes

* fixes

* fixes

* docs
2026-02-25 15:33:20 -06:00
leeguooooo e005c7251b feat(stealth): add risk-mode signals and document stealth architecture 2026-02-25 10:27:45 +09:00
leeguooooo 11eab471f1 chore(release): bump version to 0.14.0-fork.4 2026-02-25 10:10:26 +09:00
leeguooooo aa256e30c7 Merge remote-tracking branch 'upstream/main'
# Conflicts:
#	README.md
#	cli/src/connection.rs
#	cli/src/flags.rs
#	cli/src/main.rs
#	src/browser.ts
#	src/protocol.ts
2026-02-25 10:06:20 +09:00
Chris Tate c0e2b80f8c add dogfood skill for agent-driven exploratory qa (#538)
* dogfood skill

* evals

* haiku

* fixes

* caching

* fixes

* don't use npx
2026-02-24 11:35:50 -06:00
Chris Tate f319195974 add --selector flag to scroll command (#537)
* add --selector flag to scroll command

The `scroll` command uses `window.scrollBy()`, which has no effect on apps
that use custom scrollable containers (e.g. a nested div with overflow-y: auto).

The backend `handleScroll` already supports a `selector` parameter, but the CLI
never exposed it. This adds `-s` / `--selector` to the `scroll` command so users
can target a specific scrollable element:

    agent-browser scroll down 500 --selector "div.scroll-container"

Also fixes the backend to apply `direction`/`amount` when a selector is present
(previously those fields were only used in the no-selector branch).

Closes #501

* fixes
2026-02-24 07:40:46 -06:00
Chris Tate 77f2caa1bc feat: add --download-path option (#536)
* feat: add --download-path option

Adds a `--download-path` flag (and `AGENT_BROWSER_DOWNLOAD_PATH` env / `downloadPath` config key) to set a default download directory for browser downloads.

Without this, Playwright stores downloads in a temp directory that is deleted when the browser closes. The new option passes through to Playwright's `downloadsPath` on `launch()` and `launchPersistentContext()`.

Fixes #507

* improvements

* fixes

* fixes
2026-02-24 07:22:55 -06:00
leeguooooo 6f1dd39121 chore(clawhub): 改为本地 pre-push 自动同步 skill 2026-02-24 18:10:05 +09:00
leeguooooo 85d18799a4 feat(skill): 新增 agent-browser-stealth 的 OpenClaw skill 与 ClawHub 自动同步 2026-02-24 18:05:29 +09:00
leeguooooo 43e781a8d3 docs(readme): 精简文档并聚焦反爬能力 2026-02-24 18:00:19 +09:00
leeguooooo 25e8719e51 docs(readme): 补充 --delay 文本转义与 stealth 行为说明 2026-02-24 17:58:13 +09:00
leeguooooo ec011f46ff fix(stealth): 修复 launch 选项与测试基线不一致问题
在 stealth 默认策略下保留自定义 user-agent,不再被 CDP 覆盖。

同步更新 protocol/browser/launch/file-access 相关测试预期,并放宽 browser.test 的 hook 超时以消除偶发超时。
2026-02-24 17:42:12 +09:00
leeguooooo aef8fcc038 ci(release): 修复 OIDC 发布认证链路
移除 setup-node 的 registry-url 注入,避免发布步骤继承无效 NODE_AUTH_TOKEN。

发布前升级 npm 到 v11,使用独立 npmrc 并启用 provenance,以匹配 npm trusted publishing。
2026-02-24 17:24:20 +09:00
leeguooooo 96582b79fd ci(release): 调整 trusted publishing 发布流程
将 changesets/action 改为仅处理 version/PR,不再由其执行 publish。

新增发布前版本检查与独立 pnpm ci:publish 步骤,避免 OIDC 发布在 action 内失败。
2026-02-24 17:20:17 +09:00
leeguooooo 058a286326 chore(release): 发布 0.14.0-fork.3
更新 package.json 版本并写入对应 changelog 条目。
2026-02-24 17:14:06 +09:00
leeguooooo b1f27236d8 fix(cli): 修复 type/keyboard 的 --delay 参数解析
将 --delay <ms> 从输入文本中剥离并写入 delay 字段,避免搜索词混入参数。

同时支持使用 -- 终止参数解析以输入字面量 --delay 文本,并补充回归测试与帮助文档。
2026-02-24 17:13:46 +09:00
leeguooooo a5a9327b7d docs(home): 补充首页能力亮点说明
- 新增自动区域检测能力描述

- 新增验证码自动重试能力描述
2026-02-24 17:06:34 +09:00
leeguooooo 699ccbd3cb feat(cli): 强制使用用户现有浏览器并移除 profile/channel
- 禁用 --profile/AGENT_BROWSER_PROFILE 与 --channel/AGENT_BROWSER_CHANNEL,并给出项目策略提示

- 默认模式强制连接 localhost:9333,连接失败直接报错,不再自动回退新开浏览器

- 同步更新 README、技能文档、docs 与 --help 输出

- 版本升级到 0.14.0-fork.2 并同步 cli/Cargo.toml 与 Cargo.lock
2026-02-24 17:05:11 +09:00
leeguooooo 893ddfd259 feat(cdp): 默认优先连接 9333 常驻 Chrome
- 无显式连接参数时先尝试 CDP 9333,失败后回退本地浏览器启动

- 修复 CDP 页选择稳定性:过滤 omnibox 系统页、无可用页时自动创建 fallback 页

- 调整页面关闭后的 active 索引维护,降低 No page found 问题

- 新增 agent-browser-stealth 二进制入口并保持与 agent-browser 行为一致

- 同步更新 CLI 帮助、README、技能文档与 docs 说明
2026-02-24 16:22:49 +09:00
leeguooooo ea2e93dbba feat(stealth): 优化隐身对抗并引入双版本发布体系
- 将 CreepJS like headless 指标优化到 0%(headless/stealth 维持 0%)

- 新增 ActiveText 与 prefers-color-scheme 探针修复

- 版本号采用 <upstream>-fork.<fork> 格式并在 --version 输出 upstream/fork

- 更新 README、SKILL 与 docs 中的版本体系说明
2026-02-24 15:29:11 +09:00
leeguooooo 4c6afe3e69 docs(config): 移除 --stealth 配置项说明 2026-02-24 14:55:07 +09:00
leeguooooo 9f9a90cf63 docs(cli): 清理过时 stealth 环境变量说明 2026-02-24 14:54:51 +09:00
leeguooooo 0443e4ed7a feat(stealth): 默认开启并收敛 chrome 指纹特征 2026-02-24 14:54:30 +09:00
leeguooooo c5b2292caa feat(stealth): 增强浏览器级 UA 覆盖并修复背景特征 2026-02-24 14:50:03 +09:00
leeguooooo 2ed0c6f8ec feat(stealth): 进一步降低 creepjs like-headless 指标 2026-02-24 14:36:56 +09:00
leeguooooo 3a91aef4c9 feat(stealth): 优化指纹信号并同步文档与包配置 2026-02-24 14:23:59 +09:00
leeguooooo 02ebc9f328 feat(stealth): 增强指纹一致性并新增 creepjs 检测脚本 2026-02-24 14:19:25 +09:00
leeguooooo 955543b757 chore: prepare fork sync and independent release setup 2026-02-24 14:14:50 +09:00
leeguooooo ecad112707 feat(cli): 默认开启 stealth 并支持 wait 区间超时 2026-02-24 12:27:04 +09:00
leeguooooo 8932f28926 fix(stealth): 修复 headed 模式下 stealth 失效并统一策略
- 修复 launch 协议未透传 stealth 导致 --headed 下补丁失效的问题\n- 在 BrowserManager 引入 StealthPolicy,统一 local/CDP/provider 能力决策\n- 增加 launch 返回 stealth 状态并在 --debug 输出连接类型与能力\n- 补充 local/CDP 回归测试与 bot.sannysoft.com 自动检查脚本\n- 同步 README、CLI help、技能文档与 CDP 文档中的 stealth 能力矩阵
2026-02-24 12:18:47 +09:00
github-actions[bot]andgithub-actions[bot] <41898282+github-actions[bot]@users.noreply.github.com> 2fe7394dbe chore: version packages (#535)
Co-authored-by: github-actions[bot] <41898282+github-actions[bot]@users.noreply.github.com>
2026-02-23 11:03:57 -06:00
Chris Tate b7665e52b6 v0.14.0 changeset (#534)
* v0.14.0 changeset

* fixes

* improvements
2026-02-23 10:48:07 -06:00
shohuandshohu 16c4ef2da6 fix(daemon): add backpressure control and command serialization to prevent IPC EAGAIN (#529)
- Add AGENT_BROWSER_DEFAULT_TIMEOUT env var to override Playwright's
  default 60s timeout (CDP/recording 10s timeouts unaffected)
- Add backpressure-aware safeWrite() that waits for drain when socket
  buffer is full, preventing data loss under load
- Serialize command execution per socket via queue to prevent concurrent
  writes that cause buffer contention

These daemon-side fixes complement #329 (CLI-side EAGAIN retry) by
addressing the root causes: Playwright operations that outlast the
CLI's IPC timeout, and concurrent socket.write() calls that overflow
the kernel buffer.

Tested with heavy React app (1000+ DOM nodes) — 10 consecutive
snapshot commands complete without os error 35/11.

Refs #322

Co-authored-by: shohu <shohu@users.noreply.github.com>
2026-02-23 10:06:06 -06:00
ProviandClaude Opus 4.6 ad6e206a90 feat: add keyboard command for raw keyboard input (#521)
Adds `keyboard type` and `keyboard insertText` subcommands that
operate on the currently focused element without requiring a selector.

Essential for contenteditable editors (Lexical, ProseMirror, CodeMirror,
Monaco) where `type <selector>` doesn't trigger the editor's internal
event pipeline (beforeinput/DOM mutation).

- `keyboard type <text>` — page.keyboard.type() with real keystrokes
- `keyboard insertText <text>` — page.keyboard.insertText()

Note: `keyboard press` intentionally omitted — the existing top-level
`press` command already operates on current focus.

Co-authored-by: Claude Opus 4.6 <noreply@anthropic.com>
2026-02-23 09:44:08 -06:00
Lukas Malkmus f10f3f6425 cli: only warn about --annotate when explicitly passed via CLI (#531)
The warning "⚠ --annotate only applies to the screenshot command" fires
on every non-screenshot command when annotate is set in config. This is
noisy for users who set it as a persistent default.

Add cli_annotate tracking (matching the existing cli_* pattern) so the
warning only fires when --annotate is passed as a CLI flag.
2026-02-23 09:24:19 -06:00
Chris Tate c0f8f32a55 fix remote debugging (#533)
* fix remote debugging

* debug log
2026-02-23 09:19:21 -06:00
Chris Tate 12d79e4428 add --color-scheme flag for persistent dark/light mode (#528)
Fixes #519. Playwright defaults `colorScheme` to `light` on all new contexts, overriding the browser/OS dark mode setting. This is especially disruptive in CDP mode, where every reconnection resets the scheme. The `set media dark` command also didn't persist its choice to new tabs or pages.

- Add `--color-scheme <dark|light|no-preference>` flag, config key (`colorScheme`), and env var (`AGENT_BROWSER_COLOR_SCHEME`)
- Store the preference in `BrowserManager` and automatically apply it to all new contexts (via Playwright's context option) and all new pages (via `page.emulateMedia` in `setupPageTracking`)
- `set media dark/light` now also persists its choice for subsequent pages and tabs
2026-02-23 01:50:17 -06:00
Chris Tate 467b830974 fix state load failing when no browser is running (#527)
`state load` always fails with "Cannot load state while browser is running" even when no browser is running, making the command completely unusable (#526).

The daemon's auto-launch logic starts a browser before `state_load` gets to handle the command. This adds `state_load` to the exclusion list alongside `launch` and `close`, so `handleStateLoad` can perform its own launch with the state file.
2026-02-23 00:56:48 -06:00
Chris Tate 4412899379 update header/og font (#524) 2026-02-22 16:09:37 -06:00
Chris Tate fca9d7ab5d fix og (#515) 2026-02-20 08:49:53 -06:00
github-actions[bot]andgithub-actions[bot] <41898282+github-actions[bot]@users.noreply.github.com> 2b8a51b9a6 chore: version packages (#513)
Co-authored-by: github-actions[bot] <41898282+github-actions[bot]@users.noreply.github.com>
2026-02-20 00:14:32 -06:00
Chris Tate ebd87173e4 chore: add minor changeset for release (#512) 2026-02-20 00:06:52 -06:00
Chris Tate d5a667ea2d diff (#510)
* diff

* fixes

* fixes

* fixes

* fixes

* fixes

* better docs
2026-02-19 23:51:09 -06:00
github-actions[bot]andgithub-actions[bot] <41898282+github-actions[bot]@users.noreply.github.com> 9732031087 chore: version packages (#505)
Co-authored-by: github-actions[bot] <41898282+github-actions[bot]@users.noreply.github.com>
2026-02-18 22:41:17 -06:00
Chris Tate 69ffad0f04 chore: add minor changeset for release (#504) 2026-02-18 22:31:37 -06:00
Chris Tate e2e259f1e2 annotated screenshots (#503)
* screenshot annotation

* fixes

* fix CI checks

* fixes

* fixes

* fixes

* fixes

* fixes
2026-02-18 22:20:01 -06:00
github-actions[bot]andgithub-actions[bot] <41898282+github-actions[bot]@users.noreply.github.com> 06a32f4191 chore: version packages (#499)
Co-authored-by: github-actions[bot] <41898282+github-actions[bot]@users.noreply.github.com>
2026-02-18 00:40:56 -06:00
Chris Tate c6fc7df443 chore: add patch changeset for release (#498) 2026-02-18 00:34:52 -06:00
Chris Tate 98f49da196 chaining (#497) 2026-02-18 00:24:59 -06:00
github-actions[bot]andgithub-actions[bot] <41898282+github-actions[bot]@users.noreply.github.com> 85340cb432 chore: version packages (#496)
Co-authored-by: github-actions[bot] <41898282+github-actions[bot]@users.noreply.github.com>
2026-02-17 23:33:44 -06:00
Chris Tate 5dc40b4ea4 chore: add minor changeset for release (#495) 2026-02-17 23:28:34 -06:00
Andrew ImmandChris Tate 59fa36b6e2 feat: Enable capture of profiling data (#290)
* feat: Enable capture of profiling data

Adding a new set of commands:
```
agent-browser profiler start

agent-browser profiler stop trace.json
```

With this, agents can start a profiling trace, perform a set of actions, and then extract the profiling data for analysis.

**Note:** I was originally going to call it `agent-browser profile` but I realized that might cause confusion with the `--profile` flag

CDP supports a couple commands for starting/stopping a trace.
When a trace is running, it emits events that need to be picked up.
We store these locally in the daemon until the trace is completed.
When the final event is received, we dump all of them into an output file.

That file can be loaded directly into chrome devtools or another analysis tool to visualize what happened during the agentic run.

Added some basic rust tests for parsing the commands (since they have some optional / required args)

TS daemon adds ~6 tests to make sure the profiling lifecycle (including saving the output file) works as intended

* add docs

* fixes

* fixes

---------

Co-authored-by: Chris Tate <chris@ctate.dev>
2026-02-17 23:11:11 -06:00
Chris Tate 9ca182a4df add config (#494)
* add config

* improvements

* cleaner flags

* fixes

* fixes
2026-02-17 22:27:44 -06:00
Chris Tate 76df589aea update docs (#493) 2026-02-17 21:37:41 -06:00
Chris Tate 19dd2d0c0b fix(#491): auto-disable viewport for --start-maximized and --window-size args (#492)
Fixes #491

When `--start-maximized` or `--window-size` is passed as a browser arg, Playwright's default viewport (1280x720) overrides the browser's own window sizing, making those flags have no effect on the page content.

This change auto-detects those args and sets `viewport: null` so Playwright defers to the browser's window size. Explicit viewport values still take priority.

Also allows `viewport: null` in the launch protocol for agents that want to disable viewport emulation directly.
2026-02-17 20:27:38 -06:00
Chris Tate f9b33ac23d fix: reject invalid --headers JSON, empty frame commands, and --cdp + --extension combo (#488)
## Summary

- Return a `ParseError` when `--headers` receives invalid JSON instead of silently dropping the headers and proceeding
- Reject `frame` commands that provide no `selector`, `name`, or `url` (previously returned `{ switched: true }` without doing anything)
- Add missing mutual exclusion check for `--cdp` + `--extension` (extensions require a local browser, not a CDP connection)
2026-02-16 23:55:40 -06:00
Chris Tate 01efe418af fix: resolve 3 protocol bugs, improve CLI and snapshot code quality (#487)
## Summary

- Fix `allowFileAccess` being silently stripped from launch commands by adding it to the Zod schema in `protocol.ts` (the `--allow-file-access` CLI flag was not reaching the browser)
- Fix `trace stop` requiring a path argument despite help text documenting it as optional -- now works with or without a path
- Fix `addscript`/`addstyle` silently succeeding when neither `content` nor `url` is provided -- now returns a validation error
- Replace hardcoded ANSI escape code with `color::error_indicator()` in `main.rs` to respect `NO_COLOR`
- Fix double-parse pattern and add descriptive expect messages in `commands.rs`
- Fix incomplete string escaping in `snapshot.ts` `buildSelector` (use `JSON.stringify` instead of manual quote escaping)
- Simplify redundant ternary in `snapshot.ts` cursor-interactive role assignment
- Sync docs changelog with CHANGELOG.md (v0.8.1 through v0.10.0)
2026-02-16 22:47:31 -06:00
Chris Tate b7b0da5dfa docs: fix 6 documentation issues (#303, #245, #186, #134, #61, #73) (#486)
* docs: fix 6 documentation issues (#303, #245, #186, #134, #61, #73)

Addresses six open documentation issues in a single pass:

- **#303** -- Add `npx agent-browser` usage across README, SKILL.md, docs site, and `--help` output for zero-install experience. Global install is recommended as the fastest path (native Rust CLI vs Node.js indirection with npx).
- **#245** -- Document Claude Code skill installation with `npx skills add vercel-labs/agent-browser`
- **#186** -- Split installation instructions into Global (recommended), Quick Start (npx), and Project (local dependency) sections with clear guidance on when to use each
- **#134** -- Add "Why agent-browser over playwright-mcp?" comparison table to README covering output format, element selection, protocol, sessions, performance, mobile, cloud, and streaming
- **#61** -- Add "Timeouts and Slow Pages" section to SKILL.md documenting the 60s default timeout, all `wait` variants, and guidance for slow websites
- **#73** -- Replace stale `cp node_modules/...` advice with `npx skills add`, add warning against copying SKILL.md manually, add "Session Management and Cleanup" section to SKILL.md

* remove section

* fix doc
2026-02-16 22:14:43 -06:00
Giulio LeoneandCopilot d441843cca fix(#469): deduplicate cursor-interactive elements in snapshot -C (#475)
Three fixes to eliminate duplicate entries:
1. Skip elements that only inherit cursor:pointer from a parent
   (the parent element is captured instead)
2. Broaden dedup by extracting all quoted text from the ARIA tree,
   not just ref names
3. Add accepted cursor elements to the dedup set to prevent
   multiple DOM elements with the same text from duplicating

Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>
2026-02-16 11:55:08 -06:00
github-actions[bot]andgithub-actions[bot] <41898282+github-actions[bot]@users.noreply.github.com> 9cbb363190 chore: version packages (#452)
Co-authored-by: github-actions[bot] <41898282+github-actions[bot]@users.noreply.github.com>
2026-02-13 14:06:37 -06:00
Chris Tate 1112a160bd chore: add minor changeset for release (#451) 2026-02-13 13:58:54 -06:00
Aman panditandChris Tate 697b788af0 feat: add session persistence, state management commands, and --new-tab click (#184)
Rebased and fixed implementation of PR #184 features on current main:

Session persistence:
- --session-name flag and AGENT_BROWSER_SESSION_NAME env var auto-save/restore
  cookies and localStorage across browser restarts
- State files stored in ~/.agent-browser/sessions/ with owner-only permissions
- AES-256-GCM encryption via AGENT_BROWSER_ENCRYPTION_KEY env var
- Auto-expiration of old state files (AGENT_BROWSER_STATE_EXPIRE_DAYS, default 30)

State management commands:
- state list: list saved state files with metadata
- state show <file>: display state summary (cookies, origins, domains)
- state rename <old> <new>: rename state files
- state clear [name] [--all]: clear saved states
- state clean --older-than <days>: delete expired states

New --new-tab flag for click command:
- Opens link href in a new tab instead of navigating the current tab

Security hardening:
- Session name validation prevents path traversal (CLI + daemon)
- safeHeaderMerge prevents prototype pollution in header merging
- WebSocket stream server binds to 127.0.0.1 only
- State files written with 0o600 permissions

Fixes applied over the original PR:
- Use color.rs module instead of hardcoded ANSI escape codes
- Align CLI output field names with daemon response format
- Add CLI-level --session-name validation (not just daemon-side)
- Avoid adding "DOM" to tsconfig.json lib (use proper typing in evaluate)
- Keep version at 0.9.3 (matches current main)
- Centralize session name validation in daemon.ts helper
- Update all documentation (README, SKILL.md, docs site, --help output)

Co-authored-by: Chris Tate <chris@ctate.dev>
2026-02-13 11:56:20 -06:00
github-actions[bot]andgithub-actions[bot] <41898282+github-actions[bot]@users.noreply.github.com> cdd10ebb54 chore: version packages (#438)
Co-authored-by: github-actions[bot] <41898282+github-actions[bot]@users.noreply.github.com>
2026-02-13 10:59:44 -06:00
Mathias Lafeldt 323b6cdd9d Fix clippy lints (#399)
* cargo fmt

* fix: remove redundant `use libc` import (clippy::single_component_path_imports)

* fix: use `.first()` instead of `.get(0)` (clippy::get_first)

* fix: use `.copied()` instead of `.map(|s| *s)` (clippy::map_clone)

* fix: allow too_many_arguments on ensure_daemon (clippy::too_many_arguments)

* fix: use `then_some` instead of `then` with closure (clippy::unnecessary_lazy_evaluations)

* fix: use pattern match instead of redundant guard (clippy::redundant_guards)

* fix: use pattern match instead of redundant guard in commands.rs (clippy::redundant_guards)

* fix: use `contains()` instead of `iter().any()` for simple equality (clippy::manual_contains)

* Add changeset
2026-02-13 10:44:35 -06:00
Anion 604c0b9632 fix: add missing cursor field to snapshot command schema (#435)
The `-C`/`--cursor` flag was added to the CLI parser and snapshot
implementation in #374, but the Zod schema in protocol.ts was not
updated. This caused the `cursor` field to be silently stripped
during command validation, so cursor-interactive element detection
never ran.

Fixes #434
2026-02-13 08:29:13 -06:00
Chris Tate 4b776c7ba6 fix: move skill-creator out of skills/ into .agents/skills/ (#437)
- Moves `skills/skill-creator/` to `.agents/skills/skill-creator/` so that only the project-specific `agent-browser` skill remains in `skills/`
- Non-agent-browser skills like `skill-creator` are generic tooling and don't belong alongside the product skill, which was confusing to users
2026-02-13 08:26:50 -06:00
Chris Tate 9a01e8b3b5 feat: add --auto-connect flag to discover and connect to running Chrome (#432) 2026-02-12 18:37:37 -06:00
github-actions[bot]andgithub-actions[bot] <41898282+github-actions[bot]@users.noreply.github.com> 9c20979bfe chore: version packages (#430)
Co-authored-by: github-actions[bot] <41898282+github-actions[bot]@users.noreply.github.com>
2026-02-12 17:47:59 -06:00
vercel[bot]andVercel <vercel[bot]@users.noreply.github.com> 14029d2450 Add Vercel Web Analytics to Next.js (#428)
Implemented Vercel Web Analytics for Next.js (App Router)

## Summary
Successfully installed and configured @vercel/analytics package for the Next.js documentation site.

## Changes Made

### 1. Installed Dependencies
- Installed `@vercel/analytics` package using pnpm
- Command executed: `pnpm install @vercel/analytics`

### 2. Modified Files
- **docs/src/app/layout.tsx**
  - Added import: `import { Analytics } from "@vercel/analytics/next";`
  - Added `<Analytics />` component inside the `<body>` tag, right after `<SpeedInsights />`
  - Placement follows best practices for App Router projects

### 3. Updated Lock Files
- **docs/package.json** - Added @vercel/analytics to dependencies
- **docs/pnpm-lock.yaml** - Updated with new dependency tree

## Implementation Details
- This is an App Router project (uses `app/` directory structure)
- The Analytics component was added to the root layout file at `docs/src/app/layout.tsx`
- Followed the same pattern as the existing SpeedInsights component
- Preserved all existing code structure and formatting

## Verification
 Build completed successfully with no errors
 TypeScript compilation passed
 Modified file passes ESLint checks
 All 15 static pages generated correctly

## Notes
- The project already had @vercel/speed-insights installed, so the pattern for adding Analytics was consistent
- Pre-existing lint errors in mobile-nav-context.tsx and theme-toggle.tsx are unrelated to this change
- Lock files are properly updated and staged as per dependency changes

Co-authored-by: Vercel <vercel[bot]@users.noreply.github.com>
2026-02-12 17:41:49 -06:00
Chris Tate d03e238516 chore: add patch changeset for release (#429) 2026-02-12 17:41:41 -06:00
Chris Tate 221d22c14f fix: resolve stale session, ref resolution and cursor-ref collision bugs (#427) 2026-02-12 17:32:58 -06:00
vercel[bot]andVercel <vercel[bot]@users.noreply.github.com> b3b9fccd72 Add Vercel Speed Insights to Next.js (#420)
Successfully implemented Vercel Speed Insights for Next.js

## Changes Made

### 1. Installed @vercel/speed-insights package
- Used pnpm (the project's package manager) to install @vercel/speed-insights@1.3.1
- Updated package.json with the new dependency
- Updated pnpm-lock.yaml with the complete dependency tree

### 2. Integrated SpeedInsights component into root layout
- Modified: docs/src/app/layout.tsx
  - Added import: `import { SpeedInsights } from "@vercel/speed-insights/next"`
  - Added `<SpeedInsights />` component inside the `<body>` tag, placed after all other content
  - This follows the recommended pattern for Next.js 13.5+ with App Router

## Implementation Details

The project uses:
- Next.js 16.1.1 with App Router
- TypeScript
- pnpm as the package manager

The SpeedInsights component was added to the root layout (app/layout.tsx) which is the correct approach for Next.js 13.5+ projects using the App Router. The component is placed at the end of the body tag to ensure it loads after the main content.

## Verification

 Build completed successfully - no compilation errors
 All changes staged with git including the lockfile
 Package installed and integrated correctly

Note: Pre-existing linter warnings in mobile-nav-context.tsx and theme-toggle.tsx were not introduced by these changes and remain unchanged.

## Files Modified

1. docs/package.json - Added @vercel/speed-insights dependency
2. docs/pnpm-lock.yaml - Updated with new package dependencies
3. docs/src/app/layout.tsx - Added SpeedInsights import and component

The implementation follows Vercel's official documentation and best practices for Next.js App Router applications.

Co-authored-by: Vercel <vercel[bot]@users.noreply.github.com>
2026-02-12 17:28:59 -06:00
Chris Tate ec9c6a2ed9 fix: pass --executable-path to launch command in CLI (#424) 2026-02-12 13:28:25 -06:00
Chris Tate 03a8cb95d0 fix write file (#421)
* fix write file

* fix typo
2026-02-11 19:18:48 -06:00
Chris Tate 66a11aeb4c better chat (#416)
* better chat

* fixes

* fix
2026-02-11 14:17:53 -06:00
github-actions[bot]andgithub-actions[bot] <41898282+github-actions[bot]@users.noreply.github.com> ffe29b8a26 chore: version packages (#408)
Co-authored-by: github-actions[bot] <41898282+github-actions[bot]@users.noreply.github.com>
2026-02-10 14:15:12 -06:00
Chris Tate 76d23db1a9 chore: add patch changeset for release (#407) 2026-02-10 14:03:06 -06:00
Chris Tate 67cdc293f0 fix: allow localhost origins in stream server ws connections (#406) 2026-02-10 13:53:55 -06:00
Chris Tate dc53fedac0 fix: auto-switch to externally opened tabs (#404)
Update `setupContextTracking` in `BrowserManager` to auto-switch `activePageIndex` to newly opened tabs and invalidate the CDP session accordingly. This mirrors what `newTab()` and `newWindow()` already do for explicitly created tabs, and aligns CLI behavior with how real browsers focus newly opened tabs.

Fixes #384
2026-02-10 13:20:01 -06:00
Chris Tate cd4473aa64 fix: forward --exact flag to Playwright for role, label, and placeholder locators (#402) (#403)
Summary

- The `--exact` flag on `find role`, `find label`, and `find placeholder` was accepted by the CLI but silently dropped by the server. The Zod validation schema, TypeScript types, and action handlers all lacked the `exact` field, so it was stripped before reaching Playwright's `getByRole`, `getByLabel`, and `getByPlaceholder` calls.
- Added `exact` to the schema, types, and handler for all three locators so the flag is forwarded to Playwright as intended.
- Added tests confirming `exact: true` survives protocol parsing for `getbyrole`, `getbylabel`, and `getbyplaceholder`.

Fixes #402
2026-02-10 09:07:47 -06:00
Chris Tate 8e5ead85c8 fix build (#401) 2026-02-09 12:10:40 -06:00
Chris Tate e8ceafcbe1 docs: mdx, light/dark mode, ask (#400) 2026-02-09 11:16:21 -06:00
n33pm 4d8097a56f docs: add Homebrew installation instructions for macOS (#385) 2026-02-06 13:23:06 -06:00
github-actions[bot]andgithub-actions[bot] <41898282+github-actions[bot]@users.noreply.github.com> 76c30690f5 chore: version packages (#377)
Co-authored-by: github-actions[bot] <41898282+github-actions[bot]@users.noreply.github.com>
2026-02-05 00:38:29 -06:00
Chris Tate ae349451b7 chore: add patch changeset for release (#376) 2026-02-05 00:30:44 -06:00
Chris Tate 07c2372766 feat: add --allow-file-access flag for file:// URL support (#375)
* feat: add --allow-file-access flag for file:// URL support

Adds the ability to open and interact with local files using file:// URLs.
This enables use cases like viewing local PDFs, testing local HTML files,
and allowing JavaScript to access other local files via XHR.

The flag adds Chromium's --allow-file-access-from-files and --allow-file-access
launch arguments. Only supported in Chromium browsers.

Fixes #345

* fix: add cli_allow_file_access tracking to prevent spurious warning

When --allow-file-access is set via AGENT_BROWSER_ALLOW_FILE_ACCESS env var
(not CLI), don't warn about the flag being ignored when daemon is already running.
2026-02-05 00:24:28 -06:00
Chris Tate 74be667c80 feat: add cursor-interactive element detection in snapshots (#374)
* fix: only warn about ignored flags when explicitly passed via CLI

The warning about launch-time options being ignored (when daemon is
already running) was incorrectly shown when options were set via
environment variables like AGENT_BROWSER_EXECUTABLE_PATH, even when
no CLI flag was passed.

Now the warning only appears when flags are explicitly passed on the
command line, not when values come solely from environment variables.

Fixes #372

* feat: add cursor-interactive element detection in snapshots

Add -C/--cursor flag to snapshot command that detects clickable elements
that don't have proper ARIA roles but are interactive based on:
- cursor: pointer CSS style
- onclick attribute/handler
- tabindex attribute

This helps with modern web apps that use custom divs/spans as buttons.

Fixes #366

* fix: add cursor option to getSnapshot type signature
2026-02-04 23:44:58 -06:00
Chris Tate d34ce8c2d0 fix: only warn about ignored flags when explicitly passed via CLI (#373)
The warning about launch-time options being ignored (when daemon is
already running) was incorrectly shown when options were set via
environment variables like AGENT_BROWSER_EXECUTABLE_PATH, even when
no CLI flag was passed.

Now the warning only appears when flags are explicitly passed on the
command line, not when values come solely from environment variables.

Fixes #372
2026-02-04 23:14:31 -06:00
github-actions[bot]andgithub-actions[bot] <41898282+github-actions[bot]@users.noreply.github.com> 79ef5764fa chore: version packages (#360)
Co-authored-by: github-actions[bot] <41898282+github-actions[bot]@users.noreply.github.com>
2026-02-03 01:49:58 -06:00
Chris Tate 9d021bdf62 chore: add minor changeset for release (#359) 2026-02-03 01:43:58 -06:00
Chris Tate a1b992411e add iOS support (#358)
* ios

* tests

* docs

* real device

* better list

* fixes
2026-02-03 01:36:19 -06:00
github-actions[bot]andgithub-actions[bot] <41898282+github-actions[bot]@users.noreply.github.com> 3c6ae7df9d chore: version packages (#357)
Co-authored-by: github-actions[bot] <41898282+github-actions[bot]@users.noreply.github.com>
2026-02-02 21:46:53 -06:00
Chris Tate daeede49c5 chore: add patch changeset for release (#356) 2026-02-02 21:42:57 -06:00
Chris Tate 03eea8a90f fix: auto-chmod binary on first run to fix EACCES on macOS (#354)
Bun blocks postinstall scripts by default, leaving the binary without
execute permissions. The wrapper now fixes this automatically.

Fixes #344
2026-02-02 21:28:23 -06:00
github-actions[bot]andgithub-actions[bot] <41898282+github-actions[bot]@users.noreply.github.com> de859d8f6b chore: version packages (#349)
Co-authored-by: github-actions[bot] <41898282+github-actions[bot]@users.noreply.github.com>
2026-02-02 20:51:46 -06:00
Chris TateandUbuntu 17dba8f7a8 chore: add patch changeset for release (#351)
Co-authored-by: Ubuntu <ctate@ip-172-31-33-149.us-east-2.compute.internal>
2026-02-02 20:51:42 -06:00
Chris Tate 0dc36f2cff Add --stdin flag for eval command (#348)
Adds --stdin flag to read JavaScript from stdin, enabling heredoc usage
for multiline scripts without shell escaping issues.
2026-02-02 20:29:43 -06:00
github-actions[bot]andgithub-actions[bot] <41898282+github-actions[bot]@users.noreply.github.com> f770593c66 chore: version packages (#343)
Co-authored-by: github-actions[bot] <41898282+github-actions[bot]@users.noreply.github.com>
2026-02-02 19:42:32 -06:00
Chris Tate 27715884e5 chore: add patch changeset for release (#342) 2026-02-02 19:31:37 -06:00
Chris Tate e52aa49706 Add skill-creator and improve agent-browser skill (#341)
* add skills-creator

* update skill

* better docs

* minor fixes
2026-02-02 19:18:34 -06:00
Chris Tate 9c45f82193 Add base64 input for eval command (#340)
* Add base64 input for eval command

Adds -b/--base64 flag to decode script from base64, avoiding shell escaping issues for AI agents.

* Document base64 eval in SKILL.md
2026-02-02 18:52:43 -06:00
github-actions[bot]andgithub-actions[bot] <41898282+github-actions[bot]@users.noreply.github.com> bdf674a27e chore: version packages (#339)
Co-authored-by: github-actions[bot] <41898282+github-actions[bot]@users.noreply.github.com>
2026-02-02 18:26:24 -06:00
Chris Tate d24f753f51 chore: add patch changeset for release (#338) 2026-02-02 18:01:18 -06:00
Chris Tate c00dd44750 fix: improve daemon startup error handling and diagnostics (#337)
* fixes

* add debugging
2026-02-02 13:47:15 -06:00
github-actions[bot]andgithub-actions[bot] <41898282+github-actions[bot]@users.noreply.github.com> 97fd2828b5 chore: version packages (#331)
Co-authored-by: github-actions[bot] <41898282+github-actions[bot]@users.noreply.github.com>
2026-01-31 23:11:25 -06:00
Chris Tate d75350a99e chore: add patch changeset for release (#330) 2026-01-31 23:04:46 -06:00
Chris Tate 775f166bce fix: add retry logic for transient socket errors (#329)
* fix: add retry logic for transient socket errors

Fixes race condition when rapidly closing and opening browser sessions.
The daemon has a 100ms shutdown delay, which caused the CLI to detect
stale daemons as "running" and fail with EAGAIN errors.

Changes:
- Add retry logic (5 attempts, exponential backoff) for transient errors
  including EAGAIN, EOF, connection reset, and connection refused
- Add 150ms verification delay in ensure_daemon to detect shutting-down daemons
- Add cleanup_stale_files to remove leftover socket/PID files before starting
  a new daemon

Tested with 20 rapid close/open cycles and 100+ parallel commands.

* test: add unit tests for transient error detection

Extracts is_transient_error() function and adds 14 unit tests covering:
- EAGAIN errors (macOS os error 35, Linux os error 11)
- WouldBlock and Resource temporarily unavailable
- EOF and empty JSON response errors
- Connection reset (macOS os error 54, Linux os error 104)
- Broken pipe errors
- Socket not found (os error 2)
- Connection refused (macOS os error 61, Linux os error 111)
- Non-transient errors (verifies they are NOT retried)
2026-01-31 23:00:31 -06:00
github-actions[bot]andgithub-actions[bot] <41898282+github-actions[bot]@users.noreply.github.com> 32a0207ffa chore: version packages (#321)
Co-authored-by: github-actions[bot] <41898282+github-actions[bot]@users.noreply.github.com>
2026-01-29 10:52:00 -06:00
Chris Tate cb2f8c3f73 chore: add patch changeset for release (#320) 2026-01-29 10:38:04 -06:00
Chris Tate 3d24ea38fa fix: commit bin/agent-browser.js with executable permissions (#319)
Fixes #305. The file was committed with mode 644, but npm
automatically sets the executable bit on bin files, causing
git to show the file as modified after pnpm install.
2026-01-29 10:28:11 -06:00
Chris Tate 71a79f64e8 fix: sync Cargo.lock when version changes (#302)
Update sync-version.js to also run `cargo update -p agent-browser` after
updating Cargo.toml, keeping Cargo.lock in sync. Also update pre-commit
hook to stage Cargo.lock along with Cargo.toml.

This commit also brings Cargo.lock up to date (was stuck at 0.7.6).
2026-01-27 09:33:47 -06:00
github-actions[bot]andgithub-actions[bot] <41898282+github-actions[bot]@users.noreply.github.com> 72cbdc7f89 chore: version packages (#301)
Co-authored-by: github-actions[bot] <41898282+github-actions[bot]@users.noreply.github.com>
2026-01-27 09:26:52 -06:00
Chris Tate 759302ead5 v0.8.4 changeset (#300) 2026-01-27 09:20:32 -06:00
n33pm 3f74bd2171 ci(version): add version sync check between package.json and Cargo.toml (#277)
Add automated verification that package.json and cli/Cargo.toml versions
stay in sync. This prevents version drift between the npm package and
Rust CLI binary.

- Add CI job to check version sync on push/PR
- Update pre-commit hook to sync versions automatically
- Update ci:version script to include version sync step
- Add check-version-sync.js script for CI validation
2026-01-27 09:09:04 -06:00
Chris Tate 3ce441bc4e fix daemon not found (#299) 2026-01-27 09:06:09 -06:00
github-actions[bot]andgithub-actions[bot] <41898282+github-actions[bot]@users.noreply.github.com> 523d7d57f1 chore: version packages (#295)
Co-authored-by: github-actions[bot] <41898282+github-actions[bot]@users.noreply.github.com>
2026-01-27 00:19:56 -06:00
Chris Tate 4116a8ac7f chore: add patch changeset for release (#294) 2026-01-27 00:15:04 -06:00
Chris Tate 18a1abda6e test: add Windows npm global install CI test (reproduces #262) (#293)
* test: add Windows npm global install CI test (reproduces #262)

This test packs the package and installs it globally with npm,
then runs agent-browser --version. This reproduces the issue where
npm-generated shims on Windows try to invoke /bin/sh which doesn't exist.

The bin/agent-browser.js wrapper is added but not yet wired up,
so this commit should fail CI to confirm the issue.

* fix: Windows npm global install and npx support

The shell script wrapper (bin/agent-browser) with #!/bin/sh shebang
causes npm to generate Windows shims that try to invoke /bin/sh,
which doesn't exist on Windows.

This fix uses a hybrid approach:

1. Node.js wrapper (bin/agent-browser.js) as bin entry
   - Makes npx work on all platforms
   - ~100ms overhead (acceptable since npx has its own overhead)

2. postinstall patches bin entries for global installs
   - Windows: Overwrites .cmd/.ps1 shims to invoke .exe directly
   - Mac/Linux: Replaces symlink to point to native binary
   - Zero overhead for `npm i -g agent-browser` users on all platforms

Also fixes PowerShell glob expansion in CI test.

Fixes #262

* fix: Windows npm global install and npx support

The shell script wrapper (bin/agent-browser) with #!/bin/sh shebang
causes npm to generate Windows shims that try to invoke /bin/sh,
which doesn't exist on Windows.

This fix uses a hybrid approach:

1. Node.js wrapper (bin/agent-browser.js) as bin entry
   - Makes npx work on all platforms
   - ~100ms overhead (acceptable since npx has its own overhead)

2. postinstall patches bin entries for global installs
   - Windows: Overwrites .cmd/.ps1 shims to invoke .exe directly
   - Mac/Linux: Replaces symlink to point to native binary
   - Zero overhead for `npm i -g agent-browser` users on all platforms

Also adds cross-platform CI tests for npm global install to catch
regressions on all platforms (Ubuntu, macOS, Windows).

Fixes #262

* test global install

* remove dead code
2026-01-27 00:09:51 -06:00
github-actions[bot]andgithub-actions[bot] <41898282+github-actions[bot]@users.noreply.github.com> 28950b8ad2 chore: version packages (#292)
Co-authored-by: github-actions[bot] <41898282+github-actions[bot]@users.noreply.github.com>
2026-01-26 18:06:28 -06:00
Chris Tate 7e6336f65b chore: add patch changeset for release (#291) 2026-01-26 18:01:48 -06:00
Chris Tate 143c8a8f3e ci: add test for Windows CMD wrapper (#289)
* ci: add test for Windows CMD wrapper

This test will fail until the CMD wrapper is fixed to call the native binary.

* fix: Windows CMD wrapper calls native binary instead of missing index.js
2026-01-26 17:54:46 -06:00
Chris Tate 0256c8f2e9 ci: add retry logic to flaky Windows integration test (#287)
* durable windows ci

* more
2026-01-26 17:29:44 -06:00
github-actions[bot]andgithub-actions[bot] <41898282+github-actions[bot]@users.noreply.github.com> ddfaa392e4 chore: version packages (#288)
Co-authored-by: github-actions[bot] <41898282+github-actions[bot]@users.noreply.github.com>
2026-01-26 17:18:24 -06:00
Chris Tate 8eec634c6f chore: add patch changeset for release (#286) 2026-01-26 17:13:21 -06:00
Chris Tate 6a17379aaf fix: CLI binary not executable when postinstall is skipped (pnpm, bun) (#285)
* fix binary

* check binary in CI
2026-01-26 17:04:25 -06:00
Chris Tate bf5ba0a557 header (#283) 2026-01-26 14:09:17 -06:00
Chris Tate efb1923fbb v0.8.0 changelog (#282) 2026-01-26 13:46:17 -06:00
github-actions[bot]andgithub-actions[bot] <41898282+github-actions[bot]@users.noreply.github.com> 9a1cc0ed6a chore: version packages (#281)
Co-authored-by: github-actions[bot] <41898282+github-actions[bot]@users.noreply.github.com>
2026-01-26 12:03:32 -06:00
Chris Tate e0597304ec chore: add minor changeset for release (#280) 2026-01-26 11:56:51 -06:00
Li Yang e831b07f47 chore(cli): save screenshots to tmp dir when no path provided (#247)
* fix(cli): save screenshots to tmp dir when no path provided

Instead of outputting base64 to stdout (which is not useful for most CLI use cases),
screenshots without a path now save to ~/.agent-browser/tmp/screenshots/ with a
generated filename and return the path.

This makes the behavior more ergonomic for AI agents and CLI users alike.

* cleanup

* cleanup

* just revert the cargo.lock version for now

* refactor: extract getAppDir() from getSocketDir()

* docs: improve screenshot help text consistency
2026-01-26 09:08:39 -06:00
n33pm 12abdbd671 chore(cli): sync Cargo.toml version to 0.7.6 (#276) 2026-01-26 02:42:35 -06:00
Chris Tate 1b26ff886c Fix tab list command not recognizing new pages opened via clicks (#275)
## Summary

Fixed an issue where the `tab list` command couldn't recognize new pages that were opened externally (e.g., via `target="_blank"` links or popup windows). The problem occurred because context-level page tracking wasn't properly set up for all browser launch methods, causing new pages created outside of explicit `newTab()` calls to go untracked.

## Changes

- Added `setupContextTracking(context)` calls to `launch()`, `launchIncognito()`, and other context creation methods to ensure all contexts listen for new page events
- Added duplicate page checks (`!this.pages.includes(page)`) in `setupContextTracking()`, `newTab()`, and `launchIncognito()` to prevent the same page from being tracked multiple times
- Fixed `activePageIndex` calculation in `launch()` to properly set the active page index
- Enhanced comments to clarify that `setupContextTracking()` handles externally created pages (popups, new tabs from links)

## Implementation Details

The fix ensures that when a user clicks an element that opens a new tab/window, the browser context's 'page' event listener will automatically detect and track the new page. The duplicate prevention logic handles cases where both the context listener and manual page creation might try to add the same page.

Fixes #273
2026-01-26 01:25:49 -06:00
Chris Tate f862e2f7df Security: Reject cross-origin connections to daemon and stream server (#274) 2026-01-26 00:42:00 -06:00
RafaelandClaude Opus 4.5 fcee8f70d1 feat: add Kernel as cloud browser provider (#200)
Add Kernel (https://kernel.sh) as a third-party cloud browser provider,
following the same pattern as Browserbase and Browser Use integrations.

Features:
- Launch browser with `-p kernel` flag or `AGENT_BROWSER_PROVIDER=kernel`
- Configurable via environment variables:
  - KERNEL_API_KEY (required)
  - KERNEL_HEADLESS (default: false)
  - KERNEL_STEALTH (default: true)
  - KERNEL_TIMEOUT_SECONDS (default: 300)
  - KERNEL_PROFILE_NAME (optional, for persistent sessions)
- Profile find-or-create: automatically creates profile if it doesn't exist
- Profile persistence: cookies/logins saved back to profile on session close
- Uses raw fetch() calls for API communication (no SDK dependency)

Co-authored-by: Claude Opus 4.5 <noreply@anthropic.com>
2026-01-26 00:22:25 -06:00
Chris Tate a99f59cd20 Fix: check command hangs indefinitely (#272)
Fixes #257
2026-01-25 23:58:53 -06:00
Chris Tate 45506fbff0 Fix: set device does not apply deviceScaleFactor - HiDPI screenshots not possible (#270)
Fixes #255
2026-01-25 15:27:03 -06:00
shawn pana a22af0e675 generic placeholder for cloud browser provider (#260)
* docs: use generic placeholder for cloud browser provider

* docs: clarify available cloud browser providers
2026-01-25 13:55:29 -06:00
Chris Tate 79863a5180 Fix: CLI: state load / profile persistence not usable in v0.7.6 (#268)
* Fix: CLI: state load / profile persistence not usable in v0.7.6

This PR addresses issue #259

* Fix issues identified in code review
2026-01-25 13:45:17 -06:00
Chris Tate ae09fdd431 Add CLI flags for cookie URL, domain, path, httpOnly, secure, and expires (#266)
* Add CLI flags for cookie URL, domain, path, httpOnly, secure, and expires

Extends the `cookies set` command to support setting cookies with additional parameters before loading a page, solving authentication workflows where cookies need to be set for different domains.

**Key changes:**
- Added CLI flags: `--url`, `--domain`, `--path`, `--httpOnly`, `--secure`, `--sameSite`, `--expires`
- Added comprehensive test coverage for all new flags and combinations
- Updated help documentation with usage examples
- No daemon changes needed - it already supported these parameters

**Example usage:**
```bash
agent-browser cookies set session_id "abc123" --url https://app.example.com --httpOnly --secure
```

This allows setting cookies for a URL before opening the page, eliminating the need for workarounds in cross-domain authentication scenarios.

Fixes #261

* Update lock

* Fix compilation error
2026-01-25 11:53:49 -06:00
Zhiwei Li 53187a603c feat: add support for ignoring HTTPS certificate errors (#93)
* feat: add support for ignoring HTTPS certificate errors

* fix: update warning message for already running daemon to include ignore HTTPS errors option

* docs: add documentation for --ignore-https-errors option in README and SKILL.md

* feat: initialize ignore_https_errors flag in command context

* fix: change launch_cmd to mutable for cdp value handling
2026-01-24 23:54:33 -06:00
github-actions[bot]andgithub-actions[bot] <41898282+github-actions[bot]@users.noreply.github.com> 60534dfd63 chore: version packages (#243)
Co-authored-by: github-actions[bot] <41898282+github-actions[bot]@users.noreply.github.com>
2026-01-23 23:44:31 -06:00
Chris Tate a4d0c2624b chore: add patch changeset for release (#242) 2026-01-23 23:40:05 -06:00
Zach Warunek 36ea8ecb55 fix: allow null selector in screenshot command schema (#236)
The screenshot command was failing with 'Validation error: selector: Expected string, received null' when only a path was provided (e.g., 'agent-browser screenshot ~/Desktop/test.png').

The Rust CLI serializes None values as null in JSON, but the Zod schema only allowed undefined (via .optional()), not null. Changed selector field to use .nullish() which accepts both null and undefined.

Fixes issue where screenshot command without selector fails validation.
2026-01-23 17:50:11 -06:00
github-actions[bot]andgithub-actions[bot] <41898282+github-actions[bot]@users.noreply.github.com> d10fd2d545 chore: version packages (#233)
Co-authored-by: github-actions[bot] <41898282+github-actions[bot]@users.noreply.github.com>
2026-01-23 15:59:36 -06:00
Chris Tate 8c2a6ec5d2 fix: handle existing GitHub releases in workflow (#232) 2026-01-23 15:55:18 -06:00
github-actions[bot]andgithub-actions[bot] <41898282+github-actions[bot]@users.noreply.github.com> c0fd1be132 chore: version packages (#231)
Co-authored-by: github-actions[bot] <41898282+github-actions[bot]@users.noreply.github.com>
2026-01-23 15:44:44 -06:00
Chris Tate 957b5e5994 fix: ensure binary is executable after npm install (#229) 2026-01-23 15:40:46 -06:00
github-actions[bot]andgithub-actions[bot] <41898282+github-actions[bot]@users.noreply.github.com> 65d4df84ac chore: version packages (#228)
Co-authored-by: github-actions[bot] <41898282+github-actions[bot]@users.noreply.github.com>
2026-01-23 15:29:51 -06:00
Chris Tate 161d8f5c8d chore: add changeset for binary distribution fix (#227) 2026-01-23 15:25:53 -06:00
github-actions[bot]andgithub-actions[bot] <41898282+github-actions[bot]@users.noreply.github.com> f3ed1be409 chore: version packages (#226)
Co-authored-by: github-actions[bot] <41898282+github-actions[bot]@users.noreply.github.com>
2026-01-23 15:13:12 -06:00
Chris Tate 6afede28b3 chore: release v0.7.1 (#225)
Fix native binary distribution in npm package. Binaries are now built
before publishing to npm, ensuring all platforms work on installation.
2026-01-23 15:08:59 -06:00
Chris Tate 6f1c83de1b fix bin (#224) 2026-01-23 15:00:13 -06:00
Chris Tate 28129df124 fix docs (#223)
* fix: download artifacts to temp directory to avoid naming conflict

The download-artifact action creates directories named after each artifact.
When downloading to bin/, this caused conflicts because the artifact directory
names matched the binary names (e.g., bin/agent-browser-darwin-arm64/agent-browser-darwin-arm64).

Fix by downloading to artifacts/ first, then using find to move the binaries to bin/.

* fix docs
2026-01-23 13:51:23 -06:00
Chris Tate eb8325e9b4 fix: download artifacts to temp directory to avoid naming conflict (#221)
The download-artifact action creates directories named after each artifact.
When downloading to bin/, this caused conflicts because the artifact directory
names matched the binary names (e.g., bin/agent-browser-darwin-arm64/agent-browser-darwin-arm64).

Fix by downloading to artifacts/ first, then using find to move the binaries to bin/.
2026-01-23 13:15:51 -06:00
github-actions[bot]andgithub-actions[bot] <41898282+github-actions[bot]@users.noreply.github.com> 9281f46823 chore: version packages (#220)
Co-authored-by: github-actions[bot] <41898282+github-actions[bot]@users.noreply.github.com>
2026-01-23 13:06:48 -06:00
Chris Tate 316e649740 chore: add changeset for v0.7.0 release (#219) 2026-01-23 13:00:24 -06:00
Chris Tate 35d345b2b4 v0.7.0 docs (#218)
* auto-release

* fixes

* fix secret name

* update provider flag

* v0.7.0 changelog
2026-01-23 12:49:40 -06:00
Chris Tate fff4312d16 update provider flag (#217)
* auto-release

* fixes

* fix secret name

* update provider flag
2026-01-23 11:41:17 -06:00
Chris Tate 57dc7602fc auto-release (#216)
* auto-release

* fixes

* fix secret name
2026-01-23 11:19:43 -06:00
TimWhiteandChris Tate ea17db8564 fix(cli): correct output messages for state load and path-based actions (#109)
* Add files via upload

fix(cli): correct output messages for state load and path-based actions

* Add files via upload

* Update output.rs

* fix crlf

---------

Co-authored-by: Chris Tate <chris@ctate.dev>
2026-01-22 10:43:27 -06:00
Yonatan f74924cd0c feat(skills): Add hierarchical structure with references and templates (#157)
* feat(skills): Add hierarchical structure with references and templates

Adds modular documentation and executable templates to the agent-browser skill
for better AI agent consumption and progressive disclosure.

## Added

### References (deep-dive documentation)
- `references/snapshot-refs.md` - Ref lifecycle, invalidation, troubleshooting
- `references/session-management.md` - Parallel sessions, state persistence
- `references/authentication.md` - Login flows, OAuth, 2FA patterns
- `references/video-recording.md` - Recording for debugging/docs
- `references/proxy-support.md` - Proxy configuration, geo-testing

### Templates (ready-to-use workflows)
- `templates/form-automation.sh` - Form filling with validation
- `templates/authenticated-session.sh` - Login once, reuse state
- `templates/capture-workflow.sh` - Content extraction with screenshots

## Modified
- `SKILL.md` - Added reference tables linking to new documentation

## Benefits
- Progressive disclosure: Load overview first, deep dives on demand
- Reduced context: Smaller chunks for better LLM token efficiency
- Ready workflows: Copy-paste templates for common patterns

* fix(templates): Make authenticated-session.sh runnable out-of-box

Addresses review feedback: login actions were commented but verification
wasn't, causing script to fail when run as-is.

New approach:
- DISCOVERY MODE runs first (shows form structure)
- LOGIN FLOW section is fully commented as a unit
- User runs once to see refs, then customizes

┌─────────────────────────────────────────────────────────────┐
│ LOGIN FORM STRUCTURE                                        │
├─────────────────────────────────────────────────────────────┤
│ @e1 [input type="email"]                                    │
│ @e2 [input type="password"]                                 │
│ @e3 [button] "Sign In"                                      │
└─────────────────────────────────────────────────────────────┘
2026-01-22 10:26:31 -06:00
Danila PoyarkovandChris Tate 9f3c3ad933 fix(screenshot): support refs and improve error messages (#141)
* fix(screenshot): support refs and improve error messages

* fix(cli): support selector argument in screenshot command

* Fix CSS class selectors being treated as file paths

* fix(test): update screenshot test assertions

---------

Co-authored-by: Chris Tate <chris@ctate.dev>
2026-01-22 10:06:38 -06:00
Márk Magyar c046de2ec7 docs: update agent-browser skill documentation (#164) 2026-01-22 09:26:54 -06:00
55f4eaa728 feat: add download CLI commands with ref support (#183)
* feat: add download and waitfordownload CLI commands

Add CLI support for the existing download functionality in the daemon:

- `download <selector> <path>`: Click an element to trigger download
  and save to specified path
- `wait --download [path] [--timeout ms]`: Wait for any download to
  complete, optionally save to path with configurable timeout

Includes comprehensive unit tests and help documentation.

* fix: download command ref support and output message

- Fix handleDownload to use browser.getLocator() for ref selector support
- Fix CLI output to show "Downloaded to" instead of "Screenshot saved"

Co-Authored-By: Claude Opus 4.5 <noreply@anthropic.com>

---------

Co-authored-by: Claude <noreply@anthropic.com>
Co-authored-by: Chris Tate <chris@ctate.dev>
2026-01-22 08:55:19 -06:00
Chris Tate 307f970d53 fix: support WebSocket URLs in connect command (#205)
* fix: support WebSocket URLs in connect command

* address feedback
2026-01-22 08:38:53 -06:00
Lindsey SimonandChris Tate 36cca10c10 Add --profile flag for persistent browser profiles (#68)
* Add --profile flag for persistent browser profiles

Adds support for persistent browser profiles that preserve cookies,
localStorage, and login sessions across browser restarts.

Changes:
- Add --profile <path> CLI flag (flags.rs)
- Add AGENT_BROWSER_PROFILE environment variable support
- Add profile field to LaunchCommand type (types.ts)
- Use launchPersistentContext when profile is specified (browser.ts)
- Update help text and README with documentation

Usage:
  agent-browser --profile ~/.myapp-profile open myapp.com

This enables AI agents to maintain authenticated sessions across
browser restarts without re-authenticating each time.

* Expand tilde in profile path to home directory

* fix: add missing profile field to test Flags struct

---------

Co-authored-by: Chris Tate <chris@ctate.dev>
2026-01-22 08:05:25 -06:00
Tom Dale c6a92a1472 docs: add Claude Code marketplace plugin installation instructions (#181)
Document the recommended way to install the agent-browser skill using the /plugin marketplace commands introduced in PR #106.
2026-01-22 01:44:06 -06:00
Shpeedle c4f66a5922 errors doc more descriptive (#190) 2026-01-22 01:25:46 -06:00
mmhiyokoandClaude Opus 4.5 946d236d9f fix: use ~/.agent-browser for socket files instead of TMPDIR (#180)
* fix: use ~/.agent-browser for socket files instead of TMPDIR

This fixes issue #163 where different TMPDIR values (common with
tmux/screen/VSCode/IntelliJ) caused the CLI and daemon to use
different socket paths.

Socket directory priority:
1. AGENT_BROWSER_SOCKET_DIR (explicit override)
2. $XDG_RUNTIME_DIR/agent-browser (Linux standard)
3. ~/.agent-browser (fallback, like Docker Desktop)

Both CLI (Rust) and daemon (Node.js) now use the same logic.

Co-Authored-By: Claude Opus 4.5 <noreply@anthropic.com>

* fix: session list now looks in correct socket directory

- Make get_socket_dir() public in connection.rs
- Update session list to use get_socket_dir() instead of temp_dir()
- Update pid file pattern from agent-browser-{session}.pid to {session}.pid
- Add tmpdir fallback to daemon.ts when homedir is unavailable

Co-Authored-By: Claude Opus 4.5 <noreply@anthropic.com>

* test: add unit tests for socket directory resolution

Add comprehensive tests for get_socket_dir/getSocketDir to verify:
- AGENT_BROWSER_SOCKET_DIR takes priority
- Empty strings are ignored (fixes Rust/TypeScript consistency)
- XDG_RUNTIME_DIR fallback works correctly
- Home directory fallback when env vars unset

Co-Authored-By: Claude Opus 4.5 <noreply@anthropic.com>

---------

Co-authored-by: Claude Opus 4.5 <noreply@anthropic.com>
2026-01-22 01:15:25 -06:00
cb37630ccf fix: add .exe extension for Windows source binary path (#188)
The copy-native.js script was looking for 'agent-browser' but on Windows
the compiled binary is 'agent-browser.exe', causing the copy to fail.

Co-authored-by: jiazhuangai <jiazhuangai@example.com>
Co-authored-by: Claude <noreply@anthropic.com>
2026-01-22 01:11:13 -06:00
Chris Tate 61c004db94 add missing flag (#203)
* add missing flag

* clean up tests
2026-01-22 00:59:53 -06:00
OanakiajaandChris Tate 083a946aac feat: add browser launch --args, --user-agent, --proxy-bypass configuration support. (#35)
* feat: add browser launch args, user-agent, and proxy configuration support

* fix: User Agent env need added

* fix: command pass error

---------

Co-authored-by: Chris Tate <chris@ctate.dev>
2026-01-22 00:19:37 -06:00
RafaelandClaude Opus 4.5 e892bceadf feat: support remote CDP WebSocket URLs in --cdp flag (#99)
Previously, the --cdp flag only accepted a port number and connected via
http://localhost:{port}. This made it impossible to connect to remote
browser services like Kernel, Browserless, etc. that provide WebSocket URLs.

The --cdp flag now accepts either:
- A port number (e.g., 9222) for local connections
- A full WebSocket URL (e.g., wss://...) for remote browser services

Changes:
- Added cdpUrl field to LaunchCommand type
- Updated protocol validation to accept URL format with scheme validation
- Modified connectViaCDP to detect and handle both formats
- Handle numeric strings for JSON serialization edge cases
- Updated CLI to send cdpUrl or cdpPort based on input format
- Updated README with examples for remote connections

Co-authored-by: Claude Opus 4.5 <noreply@anthropic.com>
2026-01-21 21:11:14 -06:00
Aitor c4139fa389 feat: add Browser Use cloud browser as available provider (#138)
* feat: add Browser Use cloud browser
  integration

* feat: enhance Browser Use integration with provider flag support

- Updated README to reflect new usage instructions for enabling Browser Use with the `-p` flag.
- Modified CLI to parse and handle the `-p` flag for specifying the provider.
- Implemented logic in the main application to launch with the specified cloud provider.
- Adjusted BrowserManager to connect to Browser Use based on the provider flag or environment variable.
- Updated types and protocol schemas to include provider information.

* feat: add validation for mutually exclusive CLI options

- Implemented checks to prevent the use of both --cdp and --provider flags simultaneously.
- Added validation to ensure --extension cannot be used with the --provider flag.
- Enhanced error handling to provide clear feedback in both JSON and console output formats.
2026-01-21 18:01:19 -06:00
Paul KleinandKylejeong2 7123d46e7f Add Browserbase support for remote browser over CDP (#3)
* Add Browserbase support for remote browser over CDP

When BROWSERBASE_API_KEY and BROWSERBASE_PROJECT_ID env vars are set,
connect to a Browserbase session via CDP instead of launching a local browser.

* Update URLs to browserbase repo

* Add Browserbase support for remote browser over CDP

When BROWSERBASE_API_KEY and BROWSERBASE_PROJECT_ID env vars are set,
connect to a Browserbase session via CDP instead of launching a local browser.

* Update link to Browserbase Dashboard in README

* bump browserbase sdk to latest version

* remove sdk as a dep

* change name back to vercel labs

* added try catch blocks, functions to close session

* revert package names

* remove extra if statement

---------

Co-authored-by: Kylejeong2 <kylejeong21@gmail.com>
2026-01-21 17:49:21 -06:00
Chris Tate 399fd7a434 v0.6.0 changelog (#154) 2026-01-18 11:44:56 -06:00
Chris Tate 62f9b4dd6b chore: bump version to 0.6.0 (#153) 2026-01-18 11:37:26 -06:00
Chris Tate 818d9fa95e format code (#152) 2026-01-18 11:21:28 -06:00
Kye Burchard a8dcbb1222 feat: add connect command for persistent CDP sessions (#127)
Adds a `connect <port>` command that establishes a CDP connection
to a running browser. The daemon remembers the connection, so
subsequent commands work without needing --cdp on every call.

Example:
  agent-browser connect 9222
  agent-browser snapshot  # works without --cdp
  agent-browser tab
  agent-browser close
2026-01-18 10:54:12 -06:00
Mikhail Beliakovvercel[bot] <35613825+vercel[bot]@users.noreply.github.com>google-labs-jules[bot] <161369871+google-labs-jules[bot]@users.noreply.github.com>
a9fcef4579 fix: support libasound2t64 on newer Ubuntu versions (#112)
* fix: support libasound2t64 on newer Ubuntu versions

Updates the install logic to check if `libasound2t64` is available using
`apt-cache` before falling back to `libasound2`. This fixes installation
on Ubuntu 24.04 and other systems affected by the 64-bit time_t transition.

* Update cli/src/install.rs

Co-authored-by: vercel[bot] <35613825+vercel[bot]@users.noreply.github.com>

---------

Co-authored-by: google-labs-jules[bot] <161369871+google-labs-jules[bot]@users.noreply.github.com>
Co-authored-by: vercel[bot] <35613825+vercel[bot]@users.noreply.github.com>
2026-01-18 10:50:57 -06:00
Zhiwei Li 59baf97e51 fix: allow additional URL schemes in parse_command function (#125)
* fix: allow additional URL schemes in parse_command function

* fix: enhance URL validation in parse_command function to support lowercase schemes
2026-01-18 10:36:38 -06:00
0okay d02ef66c89 Refactor connection logic for Windows and hash calculationfix(cli): fix windows daemon startup and port calculation inconsistency (#79) 2026-01-18 09:14:03 -06:00
Danila Poyarkov 1689cf9eca fix(cli): handle SIGPIPE to prevent panic when piping output (#144) 2026-01-18 08:57:02 -06:00
Matthew KingandClaude Opus 4.5 03a53c9f36 feat: add Claude marketplace plugin (#106)
Co-authored-by: Claude Opus 4.5 <noreply@anthropic.com>
2026-01-18 07:49:27 -06:00
Zhiwei Li b1c0c6a366 feat: enhance response output with network request details and cleared status (#117) 2026-01-18 07:35:25 -06:00
Ryan DaigleandClaude Opus 4.5 c88734da89 feat: add NO_COLOR environment variable support (#122)
Add a centralized color module (cli/src/color.rs) that respects the
NO_COLOR environment variable per https://no-color.org/

Changes:
- Add color.rs module with helper functions for colored output
- Refactor all hardcoded ANSI escape codes to use the color module
- Add tests for color formatting functions
- Update AGENTS.md with color module usage guidelines

Co-authored-by: Claude Opus 4.5 <noreply@anthropic.com>
2026-01-17 20:37:00 -06:00
Nicenonecb 28740acecf Fix CLI/protocol mismatches for select, frame main, and headers (#45)
* fix: align CLI command payloads with protocol

* fix(cli): support multi-value select in CLI
2026-01-17 20:25:19 -06:00
jaydenfyi 4112234371 fix(cli): Output screenshot as base64 string when no path provided (#83)
* fix(cli): print screenshot base64 when no path

* chore(docs): update docs and SKILL.md

* add test for screenshot with path arg

* more minimal readme + skill change
2026-01-17 20:18:06 -06:00
Li Yang 5e08e5d077 fix: detect stale unix socket by attempting connection (#114) 2026-01-17 19:42:48 -06:00
Sanchay 42879c337a fix: respect AGENT_BROWSER_HEADED env var for headed mode (#92)
The headless option was hardcoded to true in the auto-launch section,
ignoring the AGENT_BROWSER_HEADED environment variable. This fix checks
the env var so users can run the browser in headed mode by setting
AGENT_BROWSER_HEADED=1.

Fixes #90
2026-01-17 19:22:59 -06:00
Leon Gao 412ac63b68 fix: resolve refs in input value (#139) 2026-01-17 19:05:06 -06:00
Danila Poyarkov e6e832d2bc feat: add 'get styles' command for computed styles extraction (#142) 2026-01-17 18:56:52 -06:00
Dharma b19ca760aa fix: support URL parameter in tab new command (#64)
* fix: support URL parameter in tab new command

The CLI was correctly sending the URL parameter when running
`agent-browser tab new <url>`, but the TypeScript daemon was
ignoring it because:

1. The schema didn't include the url field (stripped during validation)
2. The TabNewCommand type didn't have a url property
3. The handler didn't pass the URL to browser.newTab()
4. browser.newTab() didn't accept or use a URL parameter

This fix adds URL support throughout the chain so that
`agent-browser tab new https://example.com` now correctly
opens a new tab and navigates to the specified URL.

Fixes #62

* fix: omit url field when not provided in tab new command

Previously, the CLI always sent "url": null when no URL was provided,
which caused Zod validation to fail with "Expected string, received null".

Now the url field is only included when a URL is actually provided.

Fixes issue reported by @ctate in PR review.

* refactor: move navigation logic from BrowserManager to handleTabNew

Address review feedback:
- Add .min(1) to URL validation for consistency with navigateSchema
- Keep BrowserManager.newTab() simple (single responsibility)
- Handle navigation in handleTabNew following same pattern as handleNavigate
2026-01-17 18:53:24 -06:00
Sheingandgoogle-labs-jules[bot] <161369871+google-labs-jules[bot]@users.noreply.github.com> e7c4936bc7 fix(cli): allow null path in screenshot command validation (#101)
The Rust CLI sends `null` for the `path` argument when it is not provided,
but the Zod schema only accepted `undefined`. This change updates the
`screenshotSchema` to allow `null` values for `path`, enabling the
`screenshot` command to work without a file path argument (outputting to stdout).

Co-authored-by: google-labs-jules[bot] <161369871+google-labs-jules[bot]@users.noreply.github.com>
2026-01-17 18:41:55 -06:00
Sheinggoogle-labs-jules[bot] <161369871+google-labs-jules[bot]@users.noreply.github.com>sheing-google
7aad47d3bd fix: Prevent CDP timeout on empty URL tabs (#102)
When connecting to a browser via CDP, particularly on Android, tabs with an empty URL can cause Playwright commands to hang indefinitely. This leads to a timeout in agent-browser.

This commit fixes the issue by filtering out any pages that have an empty `page.url()` during the CDP connection process. This prevents agent-browser from attempting to interact with these problematic tabs, resolving the timeout while preserving normal pages.

Added a unit test to verify that pages with empty URLs are correctly ignored. Also increased the timeout for a flaky screencast test to improve test suite stability.

Co-authored-by: google-labs-jules[bot] <161369871+google-labs-jules[bot]@users.noreply.github.com>
Co-authored-by: sheing-google <231310897+sheing-google@users.noreply.github.com>
2026-01-17 18:16:37 -06:00
edx.eth e196ed3e35 fix(cli): align protocol action names for wheel, emulatemedia, and find locators (#143)
- mouse wheel: send 'wheel' instead of 'mousewheel'
- set media: send 'emulatemedia' instead of 'media', fix reducedMotion to be string enum
- find locators: omit 'value' field when not provided (Zod .optional() expects undefined, not null)
  - Consistently applied to: role, label, placeholder, testid, first, last, nth

Fixes #131
2026-01-17 18:11:11 -06:00
1f31452fea feat: Add video recording with Playwright native video (#116)
* feat: add video recording with Playwright native video

Adds `record start/stop` commands using Playwright's built-in video
recording. No external dependencies required (no FFmpeg).

Usage:
  agent-browser record start ./demo.webm https://example.com
  agent-browser click @e1
  agent-browser record stop

Recording creates a fresh browser context with video enabled. For smooth
demos, explore the page first to plan actions, then start recording.

Co-Authored-By: Claude Opus 4.5 <noreply@anthropic.com>

* feat: auto-capture URL and transfer state for recording

When starting a recording without a URL:
- Automatically captures current page URL
- Preserves cookies and localStorage from current session

This enables a seamless workflow:
  agent-browser open https://app.example.com
  agent-browser snapshot -i  # explore, plan
  agent-browser record start ./demo.webm  # picks up URL + auth state
  agent-browser click @e3
  agent-browser record stop

Co-Authored-By: Claude Opus 4.5 <noreply@anthropic.com>

* fix: error on non-webm recording path instead of silent coercion

Previously, specifying a non-.webm path like ./demo.mp4 would silently
change it to ./demo.webm. Now it throws a clear error telling the user
that Playwright native recording only supports WebM format.

Co-Authored-By: Claude Opus 4.5 <noreply@anthropic.com>

* fix: clean up recording temp directory after stopRecording

Previously the temp directory was created but never deleted, relying on
OS cleanup. Now we explicitly remove it after saving the video, in both
success and error paths.

Co-Authored-By: Claude Opus 4.5 <noreply@anthropic.com>

* feat: add record restart command

Adds `record restart` command that stops the current recording (if any)
and starts a new one. Also improves the error message when trying to
start recording while already recording.

Changes:
- Add restartRecording method to BrowserManager
- Add recording_restart action to protocol, types, and actions
- Add CLI parsing for `record restart <path> [url]`
- Update help text and skill documentation

Co-Authored-By: Claude Opus 4.5 <noreply@anthropic.com>

* test: add CLI tests for record restart command

Co-Authored-By: Claude Opus 4.5 <noreply@anthropic.com>

---------

Co-authored-by: Claude Opus 4.5 <noreply@anthropic.com>
Co-authored-by: Chris Tate <chris@ctate.dev>
2026-01-16 12:27:58 -06:00
NMW 3675e6bd7a feat: add --proxy flag for browser proxy configuration (#16)
* feat: add --proxy flag for browser proxy support

Add CLI flag to configure HTTP/SOCKS proxy for Playwright browser context.
Supports URL format with optional credentials: http://user:pass@host:port

* fix: improve proxy parsing error handling

- Handle malformed credentials (@ without :) by ignoring incomplete creds
- Replace unwrap() with expect() for better error messages
- Addresses Vercel bot code review suggestions

* Restaura cambios locales: soporte AGENT_BROWSER_HOME y timeout aumentado

- Agrega soporte para variable de entorno AGENT_BROWSER_HOME en connection.rs
- Aumenta timeout por defecto de 10s a 60s para conexiones más lentas

* feat: add --proxy flag for browser proxy configuration

Implements proxy support based on PR #16 with reviewer feedback:

Features:
- Parse proxy URLs: http://[user:pass@]host:port
- Support for HTTP, HTTPS, and SOCKS5 protocols
- Handle username-only auth (preserves username with empty password)
- Apply proxy to both standard and persistent contexts

Changes:
- cli/src/flags.rs: Add proxy flag parsing
- cli/src/main.rs: Add parse_proxy() with comprehensive tests
- cli/src/output.rs: Add --proxy to help output
- cli/src/commands.rs: Fix test helper to include proxy field
- src/types.ts: Add proxy to LaunchCommand interface
- src/protocol.ts: Add proxy validation schema
- src/browser.ts: Apply proxy to context creation

Tests:
- 7 unit tests for parse_proxy() covering all edge cases
- All Rust tests passing (69 tests)
- All TypeScript tests passing (168 tests)
- TypeScript typecheck passing

Resolves feedback from PR #16:
- Fixed username-only proxy handling (issue #2681046975)
- Added comprehensive unit tests
- Added --proxy to help documentation
- Used expect() instead of unwrap() for better error messages

* refactor: simplify parse_proxy function

- Remove redundant comments
- Extract server variable to reduce duplication
- Inline trivial username/password variables

All 7 proxy tests still passing.
2026-01-16 11:56:35 -06:00
Andrew GadzikandClaude Opus 4.5 fff9a146bd docs: update agent-browser skill with comprehensive command reference (#121)
Add documentation for new commands including focus, drag/drop, upload,
keydown/keyup, mouse control, cookies/storage, network interception,
tabs/windows, frames, dialogs, and browser settings.

Co-authored-by: Claude Opus 4.5 <noreply@anthropic.com>
2026-01-16 11:40:57 -06:00
edx.eth 34dcb7195a fix(cli): align console output field name with daemon response (#133)
The CLI expected a 'logs' field but the daemon returns 'messages'.
This caused 'agent-browser console' to display nothing.

Changed cli/src/output.rs to read 'messages' instead of 'logs',
matching the actual response from handleConsole in src/actions.ts.
2026-01-16 11:38:29 -06:00
Matthew KingandClaude Opus 4.5 7bdfcf8541 feat: add --version flag to CLI (#94)
Print the current version when `agent-browser --version` is passed.

Co-authored-by: Claude Opus 4.5 <noreply@anthropic.com>
2026-01-16 11:35:06 -06:00
Chris Tate 6abee37641 v0.5.0 (#78) 2026-01-13 21:54:20 -06:00
NoelandClaude Sonnet 4.5 7b43d408da fix: improve error message when element is blocked by overlay (#59)
When clicking an element that is blocked by a cookie banner or modal overlay,
the error message incorrectly showed "Element not found or not visible" even
though the element was found and visible.

The issue was in toAIFriendlyError(): the check for "Timeout" was evaluated
before "intercepts pointer events", causing the wrong error message to be
returned.

Changes:
- Reorder error detection to check "intercepts pointer events" before "Timeout"
- Improve error message to suggest dismissing modals/cookie banners
- Export toAIFriendlyError for testing
- Add focused tests for overlay blocking behavior

Before:
  Element "@e4" not found or not visible. Run 'snapshot' to see current page elements.

After:
  Element "@e4" is blocked by another element (likely a modal or overlay).
  Try dismissing any modals/cookie banners first.

Co-authored-by: Claude Sonnet 4.5 <noreply@anthropic.com>
2026-01-13 15:25:46 -06:00
Chris TateandVercel <vercel[bot]@users.noreply.github.com> 2dc093cd62 add screencast (#67)
* docs

* updates

* Fix: The handleCopy function fails to handle errors from navigator.clipboard.writeText(), causing unhandled exceptions and misleading UI feedback when clipboard operations fail.

Co-authored-by: ctate <chris@ctate.dev>

* Fix: The benchmark file uses emojis (📊, 🚀, 🔨, 📈, 📋, , ⏱️, ⚠) in console output, violating repository guidelines that forbid emojis in code and output.

Co-authored-by: ctate <chris@ctate.dev>

* Remove benchmark/run.ts from PR

* screencast

* update docs

* address comments

---------

Co-authored-by: Vercel <vercel[bot]@users.noreply.github.com>
2026-01-13 14:53:27 -06:00
Shirshak 673e2e266e feat: Add extension support (#48)
* Rebase: Add extension support

* Fix logs
2026-01-13 14:34:59 -06:00
Chris TateandVercel <vercel[bot]@users.noreply.github.com> 4713c8b520 add docs (#54)
* docs

* updates

* Fix: The handleCopy function fails to handle errors from navigator.clipboard.writeText(), causing unhandled exceptions and misleading UI feedback when clipboard operations fail.

Co-authored-by: ctate <chris@ctate.dev>

* Fix: The benchmark file uses emojis (📊, 🚀, 🔨, 📈, 📋, , ⏱️, ⚠) in console output, violating repository guidelines that forbid emojis in code and output.

Co-authored-by: ctate <chris@ctate.dev>

* Remove benchmark/run.ts from PR

---------

Co-authored-by: Vercel <vercel[bot]@users.noreply.github.com>
2026-01-13 02:54:56 -06:00
Chris Tate b4bc761168 fix builds (#55) 2026-01-13 02:41:39 -06:00
Bryan LeeandChris Tate 6eafe50952 fix incomplete build-from-source instructions (#40) (#41)
Co-authored-by: Chris Tate <chris@ctate.dev>
2026-01-13 02:21:26 -06:00
Alan JeonClaude Opus 4.5vercel[bot] <35613825+vercel[bot]@users.noreply.github.com>
95675e9d55 feat: add CDP connection support for external browsers (#24)
* feat: add CDP connection support for external browsers

Add --cdp flag to connect to browsers via Chrome DevTools Protocol.
This enables control of Electron apps, Chrome instances, or any browser
exposing a CDP endpoint.

- Add cdpPort option to launch command schema
- Implement connectViaCDP() using chromium.connectOverCDP()
- Track browser connection type for proper reconnection handling
- Collect all pages from all contexts for CDP connections

Usage: agent-browser --cdp 9222 snapshot

Co-Authored-By: Claude Opus 4.5 <noreply@anthropic.com>

* feat: enhance CDP connection handling and improve page tracking

* main.rs update

Co-authored-by: vercel[bot] <35613825+vercel[bot]@users.noreply.github.com>

* fix: verify CDP connection is alive before early return in launch()

Prevents misleading errors when the remote browser crashes by checking
isConnected() before reusing an existing browser reference.

Co-Authored-By: Claude Opus 4.5 <noreply@anthropic.com>

* fix: reconnect when CDP port changes instead of reusing existing browser

Ensures --cdp flag is respected even when a browser session already exists.
Adds tests for launch() reconnection behavior.

Co-Authored-By: Claude Opus 4.5 <noreply@anthropic.com>

* Update src/browser.ts

Co-authored-by: vercel[bot] <35613825+vercel[bot]@users.noreply.github.com>

* fix: improve CDP connection handling and validation

* feat: add CDP connection validation to ensure browser context accessibility

* Update src/browser.ts

Co-authored-by: vercel[bot] <35613825+vercel[bot]@users.noreply.github.com>

* feat: enhance CDP connection handling and add reconnect logic

* fix: improve CDP connection handling during browser closure

* fix: reset cdpPort to null during browser initialization

* feat: enhance browser launch logic to handle CDP connection switching

---------

Co-authored-by: Claude Opus 4.5 <noreply@anthropic.com>
Co-authored-by: vercel[bot] <35613825+vercel[bot]@users.noreply.github.com>
2026-01-13 00:52:47 -06:00
Byonghun Lee 97b17c98fb Update installation instructions in README (#51)
Add native build step and global link command
2026-01-13 00:39:31 -06:00
Chris Tate 57a04385c1 v0.4.4 (#32)
* v0.4.4

* 0.4.4
2026-01-12 12:12:16 -06:00
Chris Tate 1a88d7f585 custom headers via --headers (#30)
* add custom headers via --headers

* add tests

* better parsing
2026-01-12 12:01:19 -06:00
Chris Tate 4f6fd8ec5c support serverless environments (#29)
* add --executable-path

* tests

* test vercel

* fixes
2026-01-12 11:41:25 -06:00
Chris Tate 3cd0ab468f add sub --help flag (#27) 2026-01-12 10:52:31 -06:00
Chris Tate a4fcc1c198 fix windows bug (#26)
* fix windows bug

* test windows

* address feedback
2026-01-12 10:33:21 -06:00
Chris Tate 574037080c 0.4.3 (#20) 2026-01-12 01:24:48 -06:00
Chris Tate 278466764b fix readme + add missing wait flags (#19)
* fix inaccuracies

* fix wait

* address feedback
2026-01-12 01:22:10 -06:00
Chris Tate f2878c750d fix typo (#18)
* fix skill name

* 0.4.2
2026-01-12 00:58:07 -06:00
244 changed files with 119193 additions and 1121 deletions
+23
View File
@@ -0,0 +1,23 @@
# Changesets
This project uses [Changesets](https://github.com/changesets/changesets) for versioning and changelog generation.
## Adding a changeset
When you make a change that should be released, run:
```bash
pnpm changeset
```
This will prompt you to:
1. Select the type of change (patch, minor, major)
2. Write a summary of your changes
The changeset file will be committed with your PR.
## Release process
When changesets are merged to `main`, the release workflow will:
1. Create a "Version Packages" PR that updates version numbers and changelogs
2. When that PR is merged, packages are automatically published to npm
+11
View File
@@ -0,0 +1,11 @@
{
"$schema": "https://unpkg.com/@changesets/config@3.1.1/schema.json",
"changelog": "@changesets/cli/changelog",
"commit": false,
"fixed": [],
"linked": [],
"access": "public",
"baseBranch": "main",
"updateInternalDependencies": "patch",
"ignore": []
}
+19
View File
@@ -0,0 +1,19 @@
{
"$schema": "https://anthropic.com/claude-code/marketplace.schema.json",
"name": "agent-browser",
"description": "Headless browser automation for AI agents",
"owner": {
"name": "Vercel",
"email": "support@vercel.com"
},
"plugins": [
{
"name": "agent-browser",
"description": "Automates browser interactions for web testing, form filling, screenshots, and data extraction",
"source": "./",
"strict": false,
"skills": ["./skills/agent-browser"],
"category": "development"
}
]
}
+248 -15
View File
@@ -5,8 +5,19 @@ on:
branches: [main]
pull_request:
branches: [main]
workflow_dispatch:
jobs:
version-sync:
name: Version Sync Check
runs-on: ubuntu-latest
steps:
- name: Checkout repository
uses: actions/checkout@v4
- name: Check version sync
run: node scripts/check-version-sync.js
typescript:
name: TypeScript (Node ${{ matrix.node-version }})
runs-on: ubuntu-latest
@@ -45,18 +56,35 @@ jobs:
run: pnpm test
rust:
name: Rust
runs-on: ubuntu-latest
steps:
- name: Checkout repository
uses: actions/checkout@v4
- name: Setup Rust toolchain
uses: dtolnay/rust-toolchain@stable
- name: Cache Rust build artifacts
uses: Swatinem/rust-cache@v2
with:
workspaces: cli
- name: Run Rust tests
run: cargo test --profile ci --manifest-path cli/Cargo.toml
rust-cross:
name: Rust (${{ matrix.os }} - ${{ matrix.target }})
if: github.event_name != 'pull_request'
runs-on: ${{ matrix.os }}
strategy:
matrix:
include:
- os: ubuntu-latest
target: x86_64-unknown-linux-gnu
- os: macos-latest
target: aarch64-apple-darwin
- os: macos-latest
target: x86_64-apple-darwin
- os: windows-latest
- os: windows-latest-8-cores
target: x86_64-pc-windows-msvc
steps:
@@ -68,18 +96,223 @@ jobs:
with:
targets: ${{ matrix.target }}
- name: Cache Cargo dependencies
uses: actions/cache@v4
- name: Cache Rust build artifacts
uses: Swatinem/rust-cache@v2
with:
path: |
~/.cargo/bin/
~/.cargo/registry/index/
~/.cargo/registry/cache/
~/.cargo/git/db/
cli/target/
key: ${{ runner.os }}-cargo-${{ matrix.target }}-${{ hashFiles('cli/Cargo.lock') }}
restore-keys: |
${{ runner.os }}-cargo-${{ matrix.target }}-
workspaces: cli
- name: Build release binary
- name: Run Rust tests
run: cargo test --profile ci --manifest-path cli/Cargo.toml --target ${{ matrix.target }}
windows-integration:
name: Windows Integration Test
if: github.event_name != 'pull_request'
runs-on: windows-latest-8-cores
needs: rust-cross
steps:
- name: Checkout repository
uses: actions/checkout@v4
- name: Setup pnpm
uses: pnpm/action-setup@v4
with:
version: 9
- name: Setup Node.js
uses: actions/setup-node@v4
with:
node-version: 22
cache: pnpm
- name: Setup Rust toolchain
uses: dtolnay/rust-toolchain@stable
with:
targets: x86_64-pc-windows-msvc
- name: Cache Rust build artifacts
uses: Swatinem/rust-cache@v2
with:
workspaces: cli
- name: Build Rust CLI
run: cargo build --release --manifest-path cli/Cargo.toml --target x86_64-pc-windows-msvc
- name: Install npm dependencies
run: pnpm install
- name: Build TypeScript
run: pnpm build
- name: Copy CLI binary to bin directory
run: |
Copy-Item cli/target/x86_64-pc-windows-msvc/release/agent-browser.exe bin/agent-browser-win32-x64.exe
- name: Test agent-browser install command
run: |
$env:PATH = "$pwd\bin;$env:PATH"
for ($i = 1; $i -le 3; $i++) {
bin/agent-browser-win32-x64.exe install
if ($LASTEXITCODE -eq 0) { exit 0 }
Write-Host "Attempt $i failed, retrying in 10 seconds..."
Start-Sleep -Seconds 10
}
exit 1
shell: pwsh
timeout-minutes: 10
- name: Verify Chromium was installed
run: |
$playwrightPath = "$env:LOCALAPPDATA\ms-playwright"
if (Test-Path $playwrightPath) {
Write-Host "Playwright browsers installed at: $playwrightPath"
Get-ChildItem $playwrightPath -Recurse -Depth 2 | Select-Object -First 20
} else {
Write-Error "Playwright browsers not found!"
exit 1
}
shell: pwsh
- name: Test daemon lifecycle (open, snapshot, close)
run: |
$env:PATH = "$pwd\bin;$env:PATH"
Write-Host "--- Opening page ---"
bin/agent-browser-win32-x64.exe open https://example.com
if ($LASTEXITCODE -ne 0) { Write-Error "open failed"; exit 1 }
Write-Host "--- Taking snapshot ---"
$snapshot = bin/agent-browser-win32-x64.exe snapshot
if ($LASTEXITCODE -ne 0) { Write-Error "snapshot failed"; exit 1 }
Write-Host $snapshot
Write-Host "--- Closing browser ---"
bin/agent-browser-win32-x64.exe close
if ($LASTEXITCODE -ne 0) { Write-Error "close failed"; exit 1 }
Write-Host "--- Windows daemon lifecycle test passed ---"
shell: pwsh
timeout-minutes: 5
serverless-chromium:
name: Serverless Chromium (@sparticuz/chromium)
runs-on: ubuntu-latest
steps:
- name: Checkout repository
uses: actions/checkout@v4
- name: Setup pnpm
uses: pnpm/action-setup@v4
with:
version: 9
- name: Setup Node.js
uses: actions/setup-node@v4
with:
node-version: 22
cache: pnpm
- name: Install dependencies
run: pnpm install
- name: Install @sparticuz/chromium
run: pnpm add -D @sparticuz/chromium
- name: Build TypeScript
run: pnpm build
- name: Run serverless integration test
run: pnpm exec vitest run test/serverless.test.ts
global-install:
name: Global Install (${{ matrix.os }})
if: github.event_name != 'pull_request'
runs-on: ${{ matrix.os }}
needs: rust-cross
strategy:
matrix:
include:
- os: ubuntu-latest
target: x86_64-unknown-linux-gnu
binary: agent-browser-linux-x64
- os: macos-latest
target: aarch64-apple-darwin
binary: agent-browser-darwin-arm64
- os: windows-latest-8-cores
target: x86_64-pc-windows-msvc
binary: agent-browser-win32-x64.exe
steps:
- name: Checkout repository
uses: actions/checkout@v4
- name: Setup pnpm
uses: pnpm/action-setup@v4
with:
version: 9
- name: Setup Node.js
uses: actions/setup-node@v4
with:
node-version: 22
cache: pnpm
- name: Setup Rust toolchain
uses: dtolnay/rust-toolchain@stable
with:
targets: ${{ matrix.target }}
- name: Cache Rust build artifacts
uses: Swatinem/rust-cache@v2
with:
workspaces: cli
- name: Build Rust CLI
run: cargo build --release --manifest-path cli/Cargo.toml --target ${{ matrix.target }}
- name: Install npm dependencies
run: pnpm install
- name: Build TypeScript
run: pnpm build
- name: Copy CLI binary to bin directory (Unix)
if: runner.os != 'Windows'
run: cp cli/target/${{ matrix.target }}/release/agent-browser bin/${{ matrix.binary }}
- name: Copy CLI binary to bin directory (Windows)
if: runner.os == 'Windows'
run: Copy-Item cli/target/${{ matrix.target }}/release/agent-browser.exe bin/${{ matrix.binary }}
- name: Test npm global install
run: |
npm pack
npm install -g agent-browser-*.tgz
agent-browser --version
shell: bash
- name: Verify symlink points to native binary (Unix)
if: runner.os != 'Windows'
run: |
SYMLINK=$(npm prefix -g)/bin/agent-browser
TARGET=$(readlink "$SYMLINK")
echo "Symlink: $SYMLINK"
echo "Target: $TARGET"
if [[ "$TARGET" != *"${{ matrix.binary }}"* ]]; then
echo "ERROR: Symlink should point to native binary, not JS wrapper"
exit 1
fi
echo "✓ Symlink correctly points to native binary"
shell: bash
- name: Verify shim points to native binary (Windows)
if: runner.os == 'Windows'
run: |
$shimPath = "$(npm prefix -g)\agent-browser.cmd"
$content = Get-Content $shimPath -Raw
echo "Shim path: $shimPath"
echo "Shim content:"
echo $content
if ($content -notmatch "agent-browser-win32-x64\.exe") {
echo "ERROR: Shim should point to native .exe, not JS wrapper"
exit 1
}
echo "✓ Shim correctly points to native binary"
shell: pwsh
+318
View File
@@ -0,0 +1,318 @@
name: Release
on:
push:
branches:
- main
workflow_dispatch:
concurrency: ${{ github.workflow }}-${{ github.ref }}
permissions:
contents: write
pull-requests: write
id-token: write
jobs:
# Build native binaries for all platforms first
build-binaries:
name: Build ${{ matrix.name }}
runs-on: ${{ matrix.os }}
strategy:
fail-fast: false
matrix:
include:
- name: Linux x64
os: ubuntu-latest
target: x86_64-unknown-linux-gnu
binary: agent-browser-linux-x64
use_zigbuild: true
- name: Linux ARM64
os: ubuntu-latest
target: aarch64-unknown-linux-gnu
binary: agent-browser-linux-arm64
use_zigbuild: true
- name: Windows x64
os: ubuntu-latest
target: x86_64-pc-windows-gnu
binary: agent-browser-win32-x64.exe
use_zigbuild: false
- name: macOS x64
os: macos-latest
target: x86_64-apple-darwin
binary: agent-browser-darwin-x64
use_zigbuild: false
- name: macOS ARM64
os: macos-latest
target: aarch64-apple-darwin
binary: agent-browser-darwin-arm64
use_zigbuild: false
steps:
- name: Checkout repository
uses: actions/checkout@v4
- name: Setup pnpm
uses: pnpm/action-setup@v4
with:
version: 9
- name: Setup Node.js
uses: actions/setup-node@v4
with:
node-version: '22'
cache: pnpm
- name: Install npm dependencies
run: pnpm install --frozen-lockfile
- name: Sync version
run: pnpm run version:sync
- name: Setup Rust toolchain
uses: dtolnay/rust-toolchain@stable
with:
targets: ${{ matrix.target }}
- name: Install cross-compilation tools (Linux)
if: runner.os == 'Linux'
run: |
sudo apt-get update
sudo apt-get install -y gcc-aarch64-linux-gnu gcc-x86-64-linux-gnu mingw-w64
- name: Install cargo-zigbuild
if: matrix.use_zigbuild
run: |
pip3 install ziglang
cargo install cargo-zigbuild
- name: Configure Rust linkers
if: runner.os == 'Linux'
run: |
mkdir -p ~/.cargo
cat >> ~/.cargo/config.toml << 'EOF'
[target.aarch64-unknown-linux-gnu]
linker = "aarch64-linux-gnu-gcc"
[target.x86_64-pc-windows-gnu]
linker = "x86_64-w64-mingw32-gcc"
EOF
- name: Cache Rust build artifacts
uses: Swatinem/rust-cache@v2
with:
workspaces: cli
- name: Build with zigbuild
if: matrix.use_zigbuild
run: cargo zigbuild --release --manifest-path cli/Cargo.toml --target ${{ matrix.target }}
- name: Build with cargo
if: '!matrix.use_zigbuild'
run: cargo build --release --manifest-path cli/Cargo.toml --target ${{ matrix.target }}
- name: Copy binary
run: |
mkdir -p artifacts
if [[ "${{ matrix.target }}" == *"windows"* ]]; then
cp cli/target/${{ matrix.target }}/release/agent-browser.exe artifacts/${{ matrix.binary }}
else
cp cli/target/${{ matrix.target }}/release/agent-browser artifacts/${{ matrix.binary }}
chmod +x artifacts/${{ matrix.binary }}
fi
- name: Upload artifact
uses: actions/upload-artifact@v4
with:
name: ${{ matrix.binary }}
path: artifacts/${{ matrix.binary }}
retention-days: 7
# Create release PR or publish to npm (with binaries)
release:
name: Release
needs: build-binaries
runs-on: ubuntu-latest
outputs:
published: ${{ steps.publish_metadata.outputs.published }}
publishedPackages: ${{ steps.publish_metadata.outputs.publishedPackages }}
steps:
- name: Checkout Repo
uses: actions/checkout@v4
with:
fetch-depth: 0
- name: Setup pnpm
uses: pnpm/action-setup@v4
with:
version: 9
- name: Setup Node.js
uses: actions/setup-node@v4
with:
node-version: '22'
cache: pnpm
- name: Install Dependencies
run: pnpm install --frozen-lockfile
- name: Download all binary artifacts
uses: actions/download-artifact@v4
with:
path: artifacts/
- name: Move binaries to bin directory
run: |
mkdir -p bin
find artifacts -type f -name 'agent-browser-*' -exec mv {} bin/ \;
rm -rf artifacts
chmod +x bin/agent-browser-* 2>/dev/null || true
echo "Binaries in bin/:"
ls -la bin/
- name: Verify all binaries exist
run: |
EXPECTED_BINARIES=(
"agent-browser-linux-x64"
"agent-browser-linux-arm64"
"agent-browser-win32-x64.exe"
"agent-browser-darwin-x64"
"agent-browser-darwin-arm64"
)
MIN_SIZE=100000 # Binaries should be at least 100KB
ERRORS=0
for binary in "${EXPECTED_BINARIES[@]}"; do
if [ ! -f "bin/$binary" ]; then
echo "ERROR: Missing bin/$binary"
ERRORS=$((ERRORS + 1))
else
SIZE=$(stat -c%s "bin/$binary" 2>/dev/null || stat -f%z "bin/$binary")
if [ "$SIZE" -lt "$MIN_SIZE" ]; then
echo "ERROR: bin/$binary is too small ($SIZE bytes, expected >= $MIN_SIZE)"
ERRORS=$((ERRORS + 1))
else
echo "OK: bin/$binary ($SIZE bytes)"
fi
fi
done
if [ "$ERRORS" -gt 0 ]; then
echo "Error: $ERRORS binary issues found"
exit 1
fi
echo "All 5 platform binaries present and valid"
- name: Verify bundled binary versions
run: |
pnpm run verify:bundled-binaries
- name: Create Release Pull Request or Publish to npm
id: changesets
uses: changesets/action@v1
with:
version: pnpm ci:version
title: 'chore: version packages'
commit: 'chore: version packages'
env:
GITHUB_TOKEN: ${{ secrets.GITHUB_TOKEN }}
- name: Check if publish is needed
id: publish_check
if: steps.changesets.outputs.hasChangesets == 'false'
run: |
LOCAL_VERSION=$(node -p "require('./package.json').version")
REMOTE_VERSION=$(npm view agent-browser-stealth version 2>/dev/null || echo "")
echo "local_version=$LOCAL_VERSION" >> "$GITHUB_OUTPUT"
echo "remote_version=$REMOTE_VERSION" >> "$GITHUB_OUTPUT"
if [ "$LOCAL_VERSION" != "$REMOTE_VERSION" ]; then
echo "needs_publish=true" >> "$GITHUB_OUTPUT"
else
echo "needs_publish=false" >> "$GITHUB_OUTPUT"
fi
echo "Local: $LOCAL_VERSION"
echo "Remote: ${REMOTE_VERSION:-<none>}"
- name: Publish to npm (trusted publishing)
id: publish_npm
if: steps.changesets.outputs.hasChangesets == 'false' && steps.publish_check.outputs.needs_publish == 'true'
env:
NODE_AUTH_TOKEN: ''
NPM_CONFIG_USERCONFIG: /home/runner/work/_temp/trusted-npmrc
NPM_CONFIG_PROVENANCE: 'true'
run: |
npm install -g npm@^11
npm --version
printf "registry=https://registry.npmjs.org/\n" > "$NPM_CONFIG_USERCONFIG"
pnpm ci:publish
- name: Verify published npm tarball
if: steps.publish_npm.outcome == 'success'
env:
PACKAGE_NAME: agent-browser-stealth
EXPECTED_VERSION: ${{ steps.publish_check.outputs.local_version }}
run: |
pnpm run verify:registry-host-binary
- name: Set release outputs
id: publish_metadata
run: |
if [ "${{ steps.publish_npm.outcome }}" = "success" ]; then
echo "published=true" >> "$GITHUB_OUTPUT"
echo "publishedPackages=[{\"name\":\"agent-browser-stealth\",\"version\":\"${{ steps.publish_check.outputs.local_version }}\"}]" >> "$GITHUB_OUTPUT"
else
echo "published=false" >> "$GITHUB_OUTPUT"
echo "publishedPackages=[]" >> "$GITHUB_OUTPUT"
fi
# Create GitHub release with binaries after npm publish
github-release:
name: Create GitHub Release
needs: release
if: needs.release.outputs.published == 'true'
runs-on: ubuntu-latest
steps:
- name: Checkout Repo
uses: actions/checkout@v4
with:
ref: main
- name: Download all artifacts
uses: actions/download-artifact@v4
with:
path: artifacts/
- name: Move binaries to bin directory
run: |
mkdir -p bin
find artifacts -type f -name 'agent-browser-*' -exec mv {} bin/ \;
rm -rf artifacts
chmod +x bin/agent-browser-* 2>/dev/null || true
ls -la bin/
- name: Verify binaries exist
run: |
BINARY_COUNT=$(ls bin/agent-browser-* 2>/dev/null | wc -l)
if [ "$BINARY_COUNT" -lt 5 ]; then
echo "Error: Expected 5 binaries, found $BINARY_COUNT"
ls -la bin/
exit 1
fi
echo "Found $BINARY_COUNT binaries"
- name: Create GitHub Release
run: |
VERSION=$(node -p "require('./package.json').version")
TAG="v$VERSION"
# Check if release already exists
if gh release view "$TAG" &>/dev/null; then
echo "Release $TAG already exists, uploading binaries..."
gh release upload "$TAG" bin/agent-browser-* --clobber
else
echo "Creating release $TAG..."
gh release create "$TAG" \
--title "$TAG" \
--generate-notes \
bin/agent-browser-*
fi
env:
GITHUB_TOKEN: ${{ secrets.GITHUB_TOKEN }}
+17 -3
View File
@@ -4,10 +4,10 @@ node_modules/
# Build output
dist/
# Native binaries (keep the launcher scripts)
# Native binaries and build byproducts (keep JS launcher)
bin/agent-browser-*
!bin/agent-browser
!bin/agent-browser.cmd
bin/agent-browser
bin/*.d
# Rust build artifacts
cli/target/
@@ -27,10 +27,15 @@ npm-debug.log*
.DS_Store
Thumbs.db
# Python
__pycache__/
# Test artifacts
*.png
*.jpeg
*.jpg
*.webm
test/e2e/.dogfood-output/
# Package manager
package-lock.json
@@ -42,3 +47,12 @@ yarn.lock
# opensrc - source code for packages
opensrc/
# Docs site
docs/node_modules/
docs/.next/
docs/out/
docs/package-lock.json
# pnpm
.pnpm-store/
+2
View File
@@ -1 +1,3 @@
pnpm lint-staged
node scripts/sync-version.js
git add cli/Cargo.toml cli/Cargo.lock
+8
View File
@@ -0,0 +1,8 @@
if [ "${SKIP_CLAWHUB_SYNC:-0}" = "1" ]; then
echo "Skipping ClawHub sync (SKIP_CLAWHUB_SYNC=1)"
exit 0
fi
pnpm run clawhub:sync || {
echo "ClawHub sync failed. Push continues. Run 'pnpm run clawhub:sync' manually after fixing login/network."
}
+71
View File
@@ -2,9 +2,80 @@
Instructions for AI coding agents working with this codebase.
## Package Manager
This project uses **pnpm**. Always use `pnpm` instead of `npm` or `yarn` for installing dependencies, running scripts, etc. (e.g., `pnpm install`, `pnpm run build`).
## Code Style
- Do not use emojis in code, output, or documentation. Unicode symbols (✓, ✗, →, ⚠) are acceptable.
- CLI colored output uses `cli/src/color.rs`. This module respects the `NO_COLOR` environment variable. Never use hardcoded ANSI color codes.
- CLI flags must always use kebab-case (e.g., `--auto-connect`, `--allow-file-access`). Never use camelCase for flags (e.g., `--autoConnect` is wrong).
## Documentation
When adding or changing user-facing features (new flags, commands, behaviors, environment variables, etc.), update **all** of the following:
1. `cli/src/output.rs` -- `--help` output (flags list, examples, environment variables)
2. `README.md` -- Options table, relevant feature sections, examples
3. `skills/agent-browser/SKILL.md` -- so AI agents know about the feature
4. `docs/src/app/` -- the Next.js docs site (MDX pages)
5. Inline doc comments in the relevant source files
This applies to changes that either human users or AI agents would need to know about. Do not skip any of these locations.
In the `docs/src/app/` MDX files, always use HTML `<table>` syntax for tables (not markdown pipe tables). This matches the existing convention across the docs site.
## Dual Architecture (Node.js + Native)
The codebase has two daemon implementations:
- **Node.js/Playwright** (default) -- `src/daemon.ts`, `src/actions.ts`, `src/browser.ts`, and the rest of `src/`
- **Rust/Native** (experimental, `--native` or `AGENT_BROWSER_NATIVE=1`) -- `cli/src/native/daemon.rs`, `cli/src/native/actions.rs`, `cli/src/native/browser.rs`, and the rest of `cli/src/native/`
When modifying browser automation logic (commands, actions, protocol handling), changes **must** be made in **both** paths:
| Node.js Path | Native Path |
|---|---|
| `src/actions.ts` | `cli/src/native/actions.rs` |
| `src/browser.ts` | `cli/src/native/browser.rs` |
| `src/daemon.ts` | `cli/src/native/daemon.rs` |
| `src/protocol.ts` | `cli/src/native/cdp/client.rs` |
| `src/snapshot.ts` | `cli/src/native/snapshot.rs` |
| `src/state-utils.ts` | `cli/src/native/state.rs` |
New commands must be implemented in both paths, or explicitly stubbed in the native path with a clear `"Not yet implemented: {action}"` error. The goal is eventual full migration to native, but until then both paths must stay in sync.
## Testing
### Unit Tests
```bash
cd cli && cargo test
```
Runs all unit tests (~320 tests). These are fast and don't require Chrome.
### End-to-End Tests
```bash
cd cli && cargo test e2e -- --ignored --test-threads=1
```
Runs 18 e2e tests that launch real headless Chrome instances and exercise the full native daemon command pipeline. Requirements:
- Chrome must be installed
- Must run serially (`--test-threads=1`) to avoid Chrome instance contention
- Tests are `#[ignore]`'d so they don't run during normal `cargo test`
The e2e tests live in `cli/src/native/e2e_tests.rs` and cover: launch/close, navigation, snapshots, screenshots, form interaction, cookies, storage, tabs, element queries, viewport/emulation, domain filtering, diff, state management, error handling, and Phase 8 commands.
### Linting and Formatting
```bash
cd cli && cargo fmt -- --check # Check formatting
cd cli && cargo clippy # Lint
```
<!-- opensrc:start -->
+265
View File
@@ -0,0 +1,265 @@
# agent-browser
## 0.16.3-fork.1
### Patch Changes
- Sync upstream `v0.16.2` / `v0.16.3` core fixes into the fork baseline.
- Import headed-mode behavior updates from upstream.
- Improve CDP debug-port discovery by switching to `reqwest` in native Chrome probing.
- Fix dialog dismiss command parsing consistency.
- Surface daemon startup stderr on launch failure to avoid opaque timeout-only errors.
- Keep fork stealth hardening for anti-debug self-destruct flows (`disable-devtool-auto` bootstrap neutralization).
## 0.16.1-fork.5
### Patch Changes
- Harden runtime stealth against anti-debug self-destruct flows on high-risk sites:
- neutralize `disable-devtool` auto bootstrap probes by hiding the `[disable-devtool-auto]` selector entry point
- preserve normal selector behavior for non-target queries to minimize side effects
- add regression tests covering the selector patch boundary
- Expand security design docs with the anti-debug execution-plane model and clarify why page self-close/redirect is a separate surface from fingerprint scoring.
## 0.15.2-fork.0
### Patch Changes
- Merge upstream `v0.15.2` updates, including fixes for cookies clear/tab close output, daemon EPERM liveness checks, unnamed element reference matching, and docs/skills refresh.
## 0.15.1-fork.11
### Patch Changes
- Auto-attach existing browser more reliably by trying CDP localhost:9333 first, then falling back to auto-discovery before failing.
Align daemon behavior and user-facing docs/skill guidance with the same attachment policy.
## 0.15.1
### Patch Changes
- 7bd8ce9: Added support for chrome:// and chrome-extension:// URLs in navigation and recording commands. These special browser URLs are now preserved as-is instead of having https:// incorrectly prepended.
## 0.15.0
### Patch Changes
- Fix CLI typing delay parsing so `--delay` is treated as an option instead of typed text.
- Add `--delay <ms>` parsing for `type` and `keyboard type`
- Support `--` to type literal `--delay` text
- Add regression tests for parsing and delay behavior
- Update CLI help, README, skills, and docs command references
## 0.14.0
### Minor Changes
- b7665e5: - Added `keyboard` command for raw keyboard input -- type with real keystrokes, insert text, and press shortcuts at the currently focused element without needing a selector.
- Added `--color-scheme` flag and `AGENT_BROWSER_COLOR_SCHEME` env var for persistent dark/light mode preference across browser sessions.
- Fixed IPC EAGAIN errors (os error 35/11) by adding backpressure-aware socket writes, command serialization, and lowering the default Playwright timeout to 25s (configurable via `AGENT_BROWSER_DEFAULT_TIMEOUT`).
- Fixed remote debugging (CDP) reconnection.
- Fixed state load failing when no browser is running.
- Fixed `--annotate` flag warning appearing when not explicitly passed via CLI.
## 0.13.0
### Minor Changes
- ebd8717: Added new diff commands for comparing snapshots, screenshots, and URLs between page states. You can now run visual pixel diffs against baseline images, compare accessibility tree snapshots with customizable depth and selectors, and diff two URLs side-by-side with optional screenshot comparison.
## 0.12.0
### Minor Changes
- 69ffad0: Add annotated screenshots with the new --annotate flag, which overlays numbered labels on interactive elements and prints a legend mapping each label to its element ref. This enables multimodal AI models to reason about visual layout while using the same @eN refs for subsequent interactions. The flag can also be set via the AGENT_BROWSER_ANNOTATE environment variable.
## 0.11.1
### Patch Changes
- c6fc7df: Added documentation for command chaining with && across README, CLI help output, docs, and skill files, explaining how to efficiently chain multiple agent-browser commands in a single shell invocation since the browser persists via a background daemon.
## 0.11.0
### Minor Changes
- 5dc40b4: Added configuration file support with automatic loading from user and project directories, new profiler commands for Chrome DevTools profiling, computed styles getter, browser extension loading, storage state management, and iOS device emulation. Expanded click command with new-tab option, improved find command with additional actions and filtering options, and enhanced CDP connection to accept WebSocket URLs. Documentation has been significantly expanded with new sections for configuration, profiling, and proxy support.
## 0.10.0
### Minor Changes
- 1112a16: Added session persistence with automatic save/restore of cookies and localStorage across browser restarts using --session-name flag, with optional AES-256-GCM encryption for saved state data. New state management commands allow listing, showing, renaming, clearing, and cleaning up old session files. Also added --new-tab option for click commands to open links in new tabs.
## 0.9.4
### Patch Changes
- 323b6cd: Fix all Clippy lint warnings in the Rust CLI: remove redundant import, use `.first()` instead of `.get(0)`, use `.copied()` instead of `.map(|s| *s)`, use `.contains()` instead of `.iter().any()`, use `then_some` instead of lazy `then`, and simplify redundant match guards.
## 0.9.3
### Patch Changes
- d03e238: Added support for custom executable path in CLI browser launch options. Documentation site received UI improvements including a new chat component with sheet-based interface and updated dependencies.
## 0.9.2
### Patch Changes
- 76d23db: Documentation site migrated to MDX for improved content authoring, added AI-powered docs chat feature, and updated README with Homebrew installation instructions for macOS users.
## 0.9.1
### Patch Changes
- ae34945: Added --allow-file-access flag to enable opening and interacting with local file:// URLs (PDFs, HTML files) by passing Chromium flags that allow JavaScript access to local files. Added -C/--cursor flag for snapshots to include cursor-interactive elements like divs with onclick handlers or cursor:pointer styles, which is useful for modern web apps using custom clickable elements.
## 0.9.0
### Minor Changes
- 9d021bd: Add iOS Simulator and real device support for mobile Safari testing via Appium. New CLI commands include `device list` to show available simulators, `tap` and `swipe` for touch interactions, and the `--device` flag to specify which iOS device to use. Configure with `-p ios` provider flag or `AGENT_BROWSER_PROVIDER=ios` environment variable.
## 0.8.10
### Patch Changes
- 17dba8f: Add --stdin flag for eval command to read JavaScript from stdin, enabling heredoc usage for multiline scripts
- daeede4: Add --stdin flag for the eval command to read JavaScript from stdin, enabling heredoc usage for multiline scripts. Also fix binary permission issues on macOS/Linux when postinstall scripts don't run (e.g., with bun).
## 0.8.9
### Patch Changes
- 0dc36f2: Add --stdin flag for eval command to read JavaScript from stdin, enabling heredoc usage for multiline scripts
## 0.8.8
### Patch Changes
- 2771588: Added base64 encoding support for the eval command with -b/--base64 flag to avoid shell escaping issues when executing JavaScript. Updated documentation with AI agent setup instructions and reorganized the docs structure by consolidating agent mode content into the installation page.
## 0.8.7
### Patch Changes
- d24f753: Fixed browser launch options not being passed correctly when using persistent profiles, ensuring args, userAgent, proxy, and ignoreHTTPSErrors settings now work properly. Added pre-flight checks for socket path length limits and directory write permissions to provide clearer error messages when daemon startup fails. Improved error handling to properly exit with failure status when browser launch fails.
## 0.8.6
### Patch Changes
- d75350a: Improved daemon connection reliability by adding automatic retry logic for transient errors like connection resets, broken pipes, and temporary resource unavailability. The CLI now cleans up stale socket and PID files before starting a new daemon, and includes better detection of daemon responsiveness to handle race conditions during shutdown.
## 0.8.5
### Patch Changes
- cb2f8c3: Fixed version synchronization to automatically update Cargo.lock alongside Cargo.toml during releases, and made the CLI binary executable. This ensures the Rust CLI version stays in sync with the npm package version.
## 0.8.4
### Patch Changes
- 759302e: Fixed "Daemon not found" error when running through AI agents (e.g., Claude Code) by resolving symlinks in the executable path. Previously, npm global bin symlinks weren't being resolved correctly, causing intermittent daemon discovery failures.
## 0.8.3
### Patch Changes
- 4116a8a: Replaced shell-based CLI wrappers with a cross-platform Node.js wrapper to enable npx support on Windows. Added postinstall logic to patch npm's bin entry on global installs, allowing the native binary to be invoked directly with zero overhead. Added CI tests to verify global installation works correctly across all platforms.
## 0.8.2
### Patch Changes
- 7e6336f: Fixed the Windows CMD wrapper to use the native binary directly instead of routing through Node.js, improving startup performance and reliability. Added retry logic to the CI install command to handle transient failures during browser installation.
## 0.8.1
### Patch Changes
- 8eec634: Improved release workflow to validate binary file sizes and ensure binaries are executable after npm install. Updated documentation site with a new mobile navigation system and added v0.8.0 changelog entries. Reformatted CHANGELOG.md for better readability.
## v0.8.0
### New Features
- **Kernel cloud browser provider** - Connect to Kernel (https://kernel.sh) for remote browser infrastructure via `-p kernel` flag or `AGENT_BROWSER_PROVIDER=kernel`. Supports stealth mode, persistent profiles, and automatic profile find-or-create.
- **Ignore HTTPS certificate errors** - New `--ignore-https-errors` flag for working with self-signed certificates and development environments
- **Enhanced cookie management** - Extended `cookies set` command with `--url`, `--domain`, `--path`, `--httpOnly`, `--secure`, `--sameSite`, and `--expires` flags for setting cookies before page load
### Bug Fixes
- Fixed tab list command not recognizing new pages opened via clicks or `target="_blank"` links (#275)
- Fixed `check` command hanging indefinitely (#272)
- Fixed `set device` not applying deviceScaleFactor - HiDPI screenshots now work correctly (#270)
- Fixed state load and profile persistence not working in v0.7.6 (#268)
- Screenshots now save to temp directory when no path is provided (#247)
### Security
- Daemon and stream server now reject cross-origin connections (#274)
## 0.7.6
### Patch Changes
- a4d0c26: Allow null values for the screenshot selector field. Previously, passing a null selector would fail validation, but now it is properly handled as an optional value.
## 0.7.5
### Patch Changes
- 8c2a6ec: Fix GitHub release workflow to handle existing releases. If a release already exists, binaries are uploaded to it instead of failing.
## 0.7.4
### Patch Changes
- 957b5e5: Fix binary permissions on install. npm doesn't preserve execute bits, so postinstall now ensures the native binary is executable.
## 0.7.3
### Patch Changes
- 161d8f5: Fix native binary distribution in npm package. Native binaries for all platforms (Linux x64/arm64, macOS x64/arm64, Windows x64) are now correctly included when publishing.
## 0.7.2
### Patch Changes
- 6afede2: Fix native binary distribution in npm package
Native binaries for all platforms (Linux x64/arm64, macOS x64/arm64, Windows x64) are now included in the npm package. Previously, the release workflow published to npm before building binaries, causing "No binary found" errors on installation.
## 0.7.1
### Patch Changes
- Fix native binary distribution in npm package. Native binaries for all platforms (Linux x64/arm64, macOS x64/arm64, Windows x64) are now included in the npm package. Previously, the release workflow published to npm before building binaries, causing "No binary found" errors on installation.
## 0.7.0
### Minor Changes
- 316e649: ## New Features
- **Cloud browser providers** - Connect to Browserbase or Browser Use for remote browser infrastructure via `-p` flag or `AGENT_BROWSER_PROVIDER` env var
- **Persistent browser profiles** - Store cookies, localStorage, and login sessions across browser restarts with `--profile`
- **Remote CDP WebSocket URLs** - Connect to remote browser services via WebSocket URL (e.g., `--cdp "wss://..."`)
- **Download commands** - New `download` command and `wait --download` for file downloads with ref support
- **Browser launch configuration** - New `--args`, `--user-agent`, and `--proxy-bypass` flags for fine-grained browser control
- **Enhanced skills** - Hierarchical structure with references and templates for Claude Code
## Bug Fixes
- Screenshot command now supports refs and has improved error messages
- WebSocket URLs work in `connect` command
- Fixed socket file location (uses `~/.agent-browser` instead of TMPDIR)
- Windows binary path fix (.exe extension)
- State load and path-based actions now show correct output messages
## Documentation
- Added Claude Code marketplace plugin installation instructions
- Updated skill documentation with references and templates
- Improved error documentation
+227 -403
View File
@@ -1,452 +1,276 @@
# agent-browser
# agent-browser-stealth
Headless browser automation CLI for AI agents. Fast Rust CLI with Node.js fallback.
Stealth-first fork of `agent-browser` for production browser automation under anti-bot pressure.
## Installation
This README focuses on stealth architecture and principles. For full command coverage inherited from upstream, use:
### npm (recommended)
- upstream docs: <https://github.com/vercel-labs/agent-browser>
- local help: `agent-browser --help` (short alias: `abs --help`)
```bash
npm install -g agent-browser
agent-browser install # Download Chromium
```
## What This Fork Optimizes
### From Source
- Stealth is always on (legacy `launch.stealth` is accepted but ignored).
- Fingerprint surfaces are patched at multiple layers (launch args, CDP overrides, init scripts).
- Behavioral signals are humanized (typing cadence, cursor path, pacing, retry backoff).
- Region signals are auto-aligned (locale/timezone/Accept-Language) to reduce mismatch risk.
- Verification/captcha handling is policy-driven (`--risk-mode off|warn|block`).
```bash
git clone https://github.com/vercel-labs/agent-browser
cd agent-browser
pnpm install
pnpm build
agent-browser install
```
## FAQ: `agent-browser` vs `agent-browser-stealth`
### Linux Dependencies
People often ask this: "What's the anti-detection approach compared to `agent-browser-stealth` on npm?"
On Linux, install system dependencies:
- `agent-browser-stealth` on npm is the package name for this fork.
- The CLI keeps upstream-compatible command names (`agent-browser` is still the main executable, with `agent-browser-stealth` and `abs` as aliases).
- The practical difference vs upstream `agent-browser` is not one single "stealth switch"; it is a defense-in-depth stack designed for anti-bot pressure.
```bash
agent-browser install --with-deps
# or manually: npx playwright install-deps chromium
```
The core idea is layered hardening across the full automation lifecycle:
1. Connection-aware policy: choose the best available stealth capability by mode (local launch/CDP/cloud provider).
2. Fingerprint hardening: patch launch args, CDP metadata, and init-script surfaces before page code runs.
3. Behavioral humanization: non-uniform typing/mouse/wait patterns instead of perfectly mechanical actions.
4. Region coherence: auto-align locale/timezone/language signals to target geography.
5. Risk-aware control loop: detect verification/captcha signals and handle them with explicit `risk-mode` policy.
Goal: reduce detection probability and improve stability in production automation. Non-goal: "guaranteed bypass" on every target.
## Quick Start
### Install
```bash
agent-browser open example.com
agent-browser snapshot # Get accessibility tree with refs
agent-browser click @e2 # Click by ref from snapshot
agent-browser fill @e3 "test@example.com" # Fill by ref
agent-browser get text @e1 # Get text by ref
agent-browser screenshot page.png
agent-browser close
npm install -g agent-browser-stealth
agent-browser install
# same CLI, short alias
abs install
```
### Traditional Selectors (also supported)
### Minimal Usage
```bash
agent-browser click "#submit"
agent-browser fill "#email" "test@example.com"
agent-browser find role button click --name "Submit"
```
## Commands
### Core Commands
```bash
agent-browser open <url> # Navigate to URL
agent-browser click <sel> # Click element
agent-browser dblclick <sel> # Double-click element
agent-browser focus <sel> # Focus element
agent-browser type <sel> <text> # Type into element
agent-browser fill <sel> <text> # Clear and fill
agent-browser press <key> # Press key (Enter, Tab, Control+a)
agent-browser keydown <key> # Hold key down
agent-browser keyup <key> # Release key
agent-browser hover <sel> # Hover element
agent-browser select <sel> <val> # Select dropdown option
agent-browser check <sel> # Check checkbox
agent-browser uncheck <sel> # Uncheck checkbox
agent-browser scroll <dir> [px] # Scroll (up/down/left/right)
agent-browser scrollintoview <sel> # Scroll element into view
agent-browser drag <src> <tgt> # Drag and drop
agent-browser upload <sel> <files> # Upload files
agent-browser screenshot [path] # Take screenshot (--full for full page)
agent-browser pdf <path> # Save as PDF
agent-browser snapshot # Accessibility tree with refs (best for AI)
agent-browser eval <js> # Run JavaScript
agent-browser close # Close browser
```
### Get Info
```bash
agent-browser get text <sel> # Get text content
agent-browser get html <sel> # Get innerHTML
agent-browser get value <sel> # Get input value
agent-browser get attr <sel> <attr> # Get attribute
agent-browser get title # Get page title
agent-browser get url # Get current URL
agent-browser get count <sel> # Count matching elements
agent-browser get box <sel> # Get bounding box
```
### Check State
```bash
agent-browser is visible <sel> # Check if visible
agent-browser is enabled <sel> # Check if enabled
agent-browser is checked <sel> # Check if checked
```
### Find Elements (Semantic Locators)
```bash
agent-browser find role <role> <action> [value] # By ARIA role
agent-browser find text <text> <action> # By text content
agent-browser find label <label> <action> [value] # By label
agent-browser find placeholder <ph> <action> [value] # By placeholder
agent-browser find alt <text> <action> # By alt text
agent-browser find title <text> <action> # By title attr
agent-browser find testid <id> <action> [value] # By data-testid
agent-browser find first <sel> <action> [value] # First match
agent-browser find last <sel> <action> [value] # Last match
agent-browser find nth <n> <sel> <action> [value] # Nth match
```
**Actions:** `click`, `fill`, `check`, `hover`, `text`
**Examples:**
```bash
agent-browser find role button click --name "Submit"
agent-browser find text "Sign In" click
agent-browser find label "Email" fill "test@test.com"
agent-browser find first ".item" click
agent-browser find nth 2 "a" text
```
### Wait
```bash
agent-browser wait <selector> # Wait for element
agent-browser wait <ms> # Wait for time
agent-browser wait --text "Welcome" # Wait for text
agent-browser wait --url "**/dash" # Wait for URL pattern
agent-browser wait --load networkidle # Wait for load state
agent-browser wait --fn "window.ready === true" # Wait for JS condition
```
**Load states:** `load`, `domcontentloaded`, `networkidle`
### Mouse Control
```bash
agent-browser mouse move <x> <y> # Move mouse
agent-browser mouse down [button] # Press button (left/right/middle)
agent-browser mouse up [button] # Release button
agent-browser mouse wheel <dy> [dx] # Scroll wheel
```
### Browser Settings
```bash
agent-browser set viewport <w> <h> # Set viewport size
agent-browser set device <name> # Emulate device ("iPhone 14")
agent-browser set geo <lat> <lng> # Set geolocation
agent-browser set offline [on|off] # Toggle offline mode
agent-browser set headers <json> # Extra HTTP headers
agent-browser set credentials <u> <p> # HTTP basic auth
agent-browser set media [dark|light] # Emulate color scheme
```
### Cookies & Storage
```bash
agent-browser cookies # Get all cookies
agent-browser cookies set <name> <val> # Set cookie
agent-browser cookies clear # Clear cookies
agent-browser storage local # Get all localStorage
agent-browser storage local <key> # Get specific key
agent-browser storage local set <k> <v> # Set value
agent-browser storage local clear # Clear all
agent-browser storage session # Same for sessionStorage
```
### Network
```bash
agent-browser network route <url> # Intercept requests
agent-browser network route <url> --abort # Block requests
agent-browser network route <url> --body <json> # Mock response
agent-browser network unroute [url] # Remove routes
agent-browser network requests # View tracked requests
agent-browser network requests --filter api # Filter requests
```
### Tabs & Windows
```bash
agent-browser tab # List tabs
agent-browser tab new [url] # New tab (optionally with URL)
agent-browser tab <n> # Switch to tab n
agent-browser tab close [n] # Close tab
agent-browser window new # New window
```
### Frames
```bash
agent-browser frame <sel> # Switch to iframe
agent-browser frame main # Back to main frame
```
### Dialogs
```bash
agent-browser dialog accept [text] # Accept (with optional prompt text)
agent-browser dialog dismiss # Dismiss
```
### Debug
```bash
agent-browser trace start [path] # Start recording trace
agent-browser trace stop [path] # Stop and save trace
agent-browser console # View console messages
agent-browser console --clear # Clear console
agent-browser errors # View page errors
agent-browser errors --clear # Clear errors
agent-browser highlight <sel> # Highlight element
agent-browser state save <path> # Save auth state
agent-browser state load <path> # Load auth state
```
### Navigation
```bash
agent-browser back # Go back
agent-browser forward # Go forward
agent-browser reload # Reload page
```
### Setup
```bash
agent-browser install # Download Chromium browser
agent-browser install --with-deps # Also install system deps (Linux)
```
## Sessions
Run multiple isolated browser instances:
```bash
# Different sessions
agent-browser --session agent1 open site-a.com
agent-browser --session agent2 open site-b.com
# Or via environment variable
AGENT_BROWSER_SESSION=agent1 agent-browser click "#btn"
# List active sessions
agent-browser session list
# Show current session
agent-browser session
```
Each session has its own:
- Browser instance
- Cookies and storage
- Navigation history
- Authentication state
## Snapshot Options
The `snapshot` command supports filtering to reduce output size:
```bash
agent-browser snapshot # Full accessibility tree
agent-browser snapshot -i # Interactive elements only (buttons, inputs, links)
agent-browser snapshot -c # Compact (remove empty structural elements)
agent-browser snapshot -d 3 # Limit depth to 3 levels
agent-browser snapshot -s "#main" # Scope to CSS selector
agent-browser snapshot -i -c -d 5 # Combine options
```
| Option | Description |
|--------|-------------|
| `-i, --interactive` | Only show interactive elements (buttons, links, inputs) |
| `-c, --compact` | Remove empty structural elements |
| `-d, --depth <n>` | Limit tree depth |
| `-s, --selector <sel>` | Scope to CSS selector |
## Options
| Option | Description |
|--------|-------------|
| `--session <name>` | Use isolated session (or `AGENT_BROWSER_SESSION` env) |
| `--json` | JSON output (for agents) |
| `--full, -f` | Full page screenshot |
| `--name, -n` | Locator name filter |
| `--exact` | Exact text match |
| `--headed` | Show browser window (not headless) |
| `--debug` | Debug output |
## Selectors
### Refs (Recommended for AI)
Refs provide deterministic element selection from snapshots:
```bash
# 1. Get snapshot with refs
agent-browser snapshot
# Output:
# - heading "Example Domain" [ref=e1] [level=1]
# - button "Submit" [ref=e2]
# - textbox "Email" [ref=e3]
# - link "Learn more" [ref=e4]
# 2. Use refs to interact
agent-browser click @e2 # Click the button
agent-browser fill @e3 "test@example.com" # Fill the textbox
agent-browser get text @e1 # Get heading text
agent-browser hover @e4 # Hover the link
```
**Why use refs?**
- **Deterministic**: Ref points to exact element from snapshot
- **Fast**: No DOM re-query needed
- **AI-friendly**: Snapshot + ref workflow is optimal for LLMs
### CSS Selectors
```bash
agent-browser click "#id"
agent-browser click ".class"
agent-browser click "div > button"
```
### Text & XPath
```bash
agent-browser click "text=Submit"
agent-browser click "xpath=//button"
```
### Semantic Locators
```bash
agent-browser find role button click --name "Submit"
agent-browser find label "Email" fill "test@test.com"
```
## Agent Mode
Use `--json` for machine-readable output:
```bash
agent-browser snapshot --json
# Returns: {"success":true,"data":{"snapshot":"...","refs":{"e1":{"role":"heading","name":"Title"},...}}}
agent-browser get text @e1 --json
agent-browser is visible @e2 --json
```
### Optimal AI Workflow
```bash
# 1. Navigate and get snapshot
agent-browser open example.com
agent-browser snapshot -i --json # AI parses tree and refs
# 2. AI identifies target refs from snapshot
# 3. Execute actions using refs
agent-browser open https://example.com
agent-browser snapshot -i
agent-browser click @e2
agent-browser fill @e3 "input text"
# 4. Get new snapshot if page changed
agent-browser snapshot -i --json
```
## Headed Mode
Show the browser window for debugging:
### Default: Auto Group Agent Tabs (CDP + Plugin)
```bash
agent-browser open example.com --headed
agent-browser open https://example.com
# In CDP mode, tabs are grouped when the tab-group extension is installed
# Override group title
agent-browser --tab-group "My Agent Group" open https://example.com
```
This opens a visible browser window instead of running headless.
- CDP (`--cdp` / `--auto-connect`) keeps working unchanged.
- If the extension is installed and handshake succeeds, agent tabs are grouped by session:
- session=`default`: `Agent Browser Stealth`
- other sessions: `Agent Browser Stealth • <session>`
- If the extension is missing/unavailable, commands continue normally with silent no-op (no warning/error unless `AGENT_BROWSER_DEBUG=1`).
- Env overrides:
- `AGENT_BROWSER_TAB_GROUP` for base title
- `AGENT_BROWSER_TAB_GROUP_PLUGIN_ID` for expected extension ID
## Architecture
Install once in Chrome: load unpacked extension from `extensions/tab-group-cdp/` (extension name: `agent-browser-stealth`).
agent-browser uses a client-daemon architecture:
### Extension Capabilities (`agent-browser-stealth`)
1. **Rust CLI** (fast native binary) - Parses commands, communicates with daemon
2. **Node.js Daemon** - Manages Playwright browser instance
3. **Fallback** - If native binary unavailable, uses Node.js directly
- Session window isolation: tabs are kept in their session window when possible.
- Configurable isolation controls: side panel can toggle `strictWindowIsolation` and cross-window activation guard.
- Session-aware grouping: deterministic group color, default session expanded, non-default sessions collapsed.
- Download archive routing: downloads from managed tabs are routed to `agent-browser-stealth/<session>/...`.
- Domain allowlist fallback: when allowlist is configured for a session, extension can force-block out-of-policy tabs to `about:blank`.
- Risk hints (debug only): suspicious host/TLD hints are returned via handshake and printed only when `AGENT_BROWSER_DEBUG=1`.
- Side panel browser controls: open/back/forward/reload, click/fill/press by CSS selector, run shortcut commands, and switch/close tabs.
- Side panel developer signals: capture page console errors/warnings, fetch/xhr network events, command history, and live DOM snapshots.
- Workflow automation: record actions into workflows, run workflows, map workflows to slash shortcuts, and schedule runs (daily/weekly/monthly/yearly).
- Side panel operations console: view session/tab/group mapping, focus a session, keep only one session, clean empty groups, edit session allowlist, and toggle auto-clean.
The daemon starts automatically on first command and persists between commands for fast subsequent operations.
## Stealth Architecture
## Platforms
| Platform | Binary | Fallback |
|----------|--------|----------|
| macOS ARM64 | ✅ Native Rust | Node.js |
| macOS x64 | ✅ Native Rust | Node.js |
| Linux ARM64 | ✅ Native Rust | Node.js |
| Linux x64 | ✅ Native Rust | Node.js |
| Windows | - | Node.js |
## Usage with AI Agents
### Just ask the agent
The simplest approach - just tell your agent to use it:
```
Use agent-browser to test the login flow. Run agent-browser --help to see available commands.
```mermaid
flowchart TD
A["Command Input"] --> B["Stealth Policy Resolver"]
B --> C["Connection Mode Detection"]
C --> D["Launch Layer: Chromium Args"]
C --> E["CDP Layer: UA + Metadata Override"]
C --> F["Context Layer: Init Script Patches"]
D --> G["Behavior Layer: Humanized Interaction"]
E --> G
F --> G
G --> H["Risk Layer: Verification Detection and Handling"]
H --> I["Response with warnings and riskSignals"]
```
The `--help` output is comprehensive and most agents can figure it out from there.
### Policy by Connection Mode
### AGENTS.md / CLAUDE.md
| Mode | Stealth Capabilities | Notes |
| --------------------------------------- | ------------------------------------------------------------- | -------------------------------------------- |
| Local Chromium launch | Chromium launch args + CDP UA override + context init scripts | Most complete stack |
| Existing browser via CDP | CDP UA override + context init scripts | No local Chromium arg injection |
| Cloud provider (browserbase/browseruse) | Context init scripts | Remote browser runtime controls launch layer |
| Kernel provider | Context init scripts + provider-managed stealth | Provider-side stealth may also apply |
For more consistent results, add to your project or global instructions file:
## Principle 1: Always-On Stealth with Explicit Boundaries
```markdown
## Browser Automation
- Stealth defaults to enabled and does not depend on a runtime toggle.
- Project policy forbids:
- `--profile` / `AGENT_BROWSER_PROFILE`
- `--channel` / `AGENT_BROWSER_CHANNEL`
- Default CLI policy auto-attaches an existing browser: try CDP `localhost:9333` first, then auto-discovery unless explicit connection options are provided.
Use `agent-browser` for web automation. Run `agent-browser --help` for all commands.
## Principle 2: Multi-Layer Fingerprint Hardening
Core workflow:
1. `agent-browser open <url>` - Navigate to page
2. `agent-browser snapshot -i` - Get interactive elements with refs (@e1, @e2)
3. `agent-browser click @e1` / `fill @e2 "text"` - Interact using refs
4. Re-snapshot after page changes
```
### 2.1 Launch Layer (Local Chromium)
### Claude Code Skill
Injected Chromium args:
For Claude Code, a [skill](https://platform.claude.com/docs/en/agents-and-tools/agent-skills/best-practices) provides richer context:
- `--disable-blink-features=AutomationControlled`
- `--use-gl=angle`
- `--use-angle=default`
If no custom UA is set, the runtime UA is normalized to remove `HeadlessChrome` tokens.
### 2.2 CDP Layer (Browser/Page Targets)
- Uses `Emulation.setUserAgentOverride` to align:
- `userAgent`
- `acceptLanguage`
- `userAgentMetadata` brands and versions
- Applies overrides for existing/new targets, including worker-relevant contexts.
- Forces opaque white background (`Emulation.setDefaultBackgroundColorOverride`) to avoid headless transparency fingerprints.
### 2.3 Context Init-Script Layer (Patch Inventory)
The init script patch set is injected before page scripts and currently includes:
1. `navigator.webdriver` removal (including prototype-level cleanup).
2. CSS webdriver heuristic neutralization (`CSS.supports('border-end-end-radius: initial')` probe).
3. `window.chrome.runtime` bootstrap for missing runtime surfaces.
4. Locale/language normalization (`navigator.language`, `navigator.languages`).
5. Realistic `navigator.plugins` and `navigator.mimeTypes`.
6. `navigator.permissions.query` normalization for notifications.
7. WebGL vendor/renderer masking when SwiftShader indicators are present.
8. `cdc_` property cleanup on document/documentElement.
9. Window/screen dimension normalization (`outerWidth/outerHeight/screenX/screenY`).
10. Screen availability patching (`availWidth/availHeight`).
11. Hardware concurrency stabilization.
12. Notification permission consistency.
13. Active text color heuristic patching.
14. `navigator.connection` normalization.
15. Worker network signal normalization (`downlinkMax`).
16. `prefers-color-scheme` light-mode heuristic neutralization.
17. `navigator.share` exposure.
18. `navigator.contacts` exposure.
19. `contentIndex` exposure.
20. `navigator.pdfViewerEnabled` normalization.
21. Media devices surface normalization.
22. `navigator.userAgent` cleanup (strip `HeadlessChrome`).
23. `navigator.userAgentData` brand cleanup.
24. `performance.memory` stabilization.
25. Default background color patching at script level.
## Principle 3: Behavioral Humanization
- Navigation pacing jitter before `goto` (short randomized delay).
- Typing jitter for `type --delay` and `keyboard type --delay`:
- per-character randomized delay around the requested base delay (about ±40%).
- Click path humanization:
- cursor moves on a Bezier-like curve before click.
- Wait supports random ranges (`wait min-max`) for non-uniform timing.
## Principle 4: Region Signal Alignment
Before navigation, the runtime derives region hints from target URL TLD and aligns:
- locale
- timezone
- `Accept-Language`
Examples of built-in mappings include `tw`, `jp`, `kr`, `sg`, `de`, `fr`, `uk`, `in`, `au`.
Manual overrides are supported:
- `AGENT_BROWSER_LOCALE`
- `AGENT_BROWSER_TIMEZONE` (or `TZ`)
## Principle 5: Verification-Aware Risk Control
When a navigation lands on verification/captcha pages, structured risk signals are generated from URL/title/page-text evidence.
`riskSignals` include:
- `code`
- `source` (`url` or `title`)
- `evidence`
- `confidence`
### Risk Mode
- `warn` (default): wait for auto-clear, then retry with randomized backoff and return warnings + `riskSignals`.
- `block`: fail fast once verification/captcha interstitial is detected.
- `off`: skip detection/retry path.
```bash
cp -r node_modules/agent-browser/skills/agent-browser .claude/skills/
agent-browser --risk-mode warn open https://example.com
agent-browser --risk-mode block open https://example.com
AGENT_BROWSER_RISK_MODE=off agent-browser open https://example.com
```
Or download:
```mermaid
flowchart TD
A["Navigate"] --> B["Collect URL/Title/Text Signals"]
B --> C{"risk-mode"}
C -->|off| D["Return Success"]
C -->|block| E["Return Error with First Signal"]
C -->|warn| F["Wait for auto-clear, then retry up to 2 times"]
F --> G{"Signals Cleared"}
G -->|yes| H["Return Success + recovery warning + riskSignals"]
G -->|no| I["Return Success + warning + riskSignals"]
```
## Operational Recommendations
- Prefer `--headed` for high-friction targets.
- Reuse session state with one stable `--session-name` for continuity (when omitted, it defaults to `default`).
- Keep locale/timezone consistent with target market.
- For challenge-heavy pages, prefer `--wait-until domcontentloaded` on `open`/`navigate` to avoid `load` stalls.
- Use `--risk-mode block` in strict pipelines that require explicit operator intervention on verification pages.
- For `cookies set`, use either `--url <url>`, or `--domain <domain> --path <path>` together.
- If `--url`, `--domain`, and `--path` are all omitted, the cookie is scoped from the current page URL.
## Validation Scripts
Run public detector checks after stealth changes:
```bash
mkdir -p .claude/skills/agent-browser
curl -o .claude/skills/agent-browser/SKILL.md \
https://raw.githubusercontent.com/vercel-labs/agent-browser/main/skills/agent-browser/SKILL.md
node scripts/check-sannysoft-webdriver.js --binary ./cli/target/release/agent-browser
node scripts/check-creepjs-headless.js --binary ./cli/target/release/agent-browser
node scripts/check-stealth-regression.js --binary ./cli/target/release/agent-browser
pnpm run check:turnstile-testkey
```
## Doctor Diagnostics
Use `doctor` to quickly diagnose local CDP, sourceURL sanitization, and tab-group plugin readiness:
```bash
agent-browser doctor
agent-browser --json doctor
```
`doctor` checks:
- CDP probe status (preferred `:9333` plus common ports)
- DevToolsActivePort discovery from local Chrome profiles
- CDP Runtime.evaluate sourceURL sanitization probe
- Plugin handshake page context check (internal page vs normal `http(s)` page)
- Tab-group extension handshake (when currently attached in CDP mode)
## Upstream Compatibility
This fork intentionally keeps command workflows close to upstream while concentrating custom behavior in stealth, policy, and anti-detection handling.
## License
Apache-2.0
-26
View File
@@ -1,26 +0,0 @@
#!/bin/sh
# agent-browser CLI wrapper
# Detects OS/arch and runs the appropriate native binary
SCRIPT="$0"
while [ -L "$SCRIPT" ]; do
SCRIPT_DIR="$(cd "$(dirname "$SCRIPT")" && pwd)"
SCRIPT="$(readlink "$SCRIPT")"
case "$SCRIPT" in /*) ;; *) SCRIPT="$SCRIPT_DIR/$SCRIPT" ;; esac
done
SCRIPT_DIR="$(cd "$(dirname "$SCRIPT")" && pwd)"
OS=$(uname -s | tr '[:upper:]' '[:lower:]')
ARCH=$(uname -m)
case "$OS" in darwin) OS="darwin" ;; linux) OS="linux" ;; mingw*|msys*|cygwin*) OS="win32" ;; esac
case "$ARCH" in x86_64|amd64) ARCH="x64" ;; aarch64|arm64) ARCH="arm64" ;; esac
BINARY="$SCRIPT_DIR/agent-browser-${OS}-${ARCH}"
if [ -f "$BINARY" ] && [ -x "$BINARY" ]; then
exec "$BINARY" "$@"
fi
echo "Error: No binary found for ${OS}-${ARCH}" >&2
echo "Run 'npm run build:native' to build for your platform" >&2
exit 1
-5
View File
@@ -1,5 +0,0 @@
@echo off
setlocal
set "SCRIPT_DIR=%~dp0"
node "%SCRIPT_DIR%..\dist\index.js" %*
exit /b %errorlevel%
+109
View File
@@ -0,0 +1,109 @@
#!/usr/bin/env node
/**
* Cross-platform CLI wrapper for agent-browser
*
* This wrapper enables npx support on Windows where shell scripts don't work.
* For global installs, postinstall.js patches the shims to invoke the native
* binary directly (zero overhead).
*/
import { spawn } from 'child_process';
import { existsSync, accessSync, chmodSync, constants } from 'fs';
import { dirname, join } from 'path';
import { fileURLToPath } from 'url';
import { platform, arch } from 'os';
const __dirname = dirname(fileURLToPath(import.meta.url));
// Map Node.js platform/arch to binary naming convention
function getBinaryName() {
const os = platform();
const cpuArch = arch();
let osKey;
switch (os) {
case 'darwin':
osKey = 'darwin';
break;
case 'linux':
osKey = 'linux';
break;
case 'win32':
osKey = 'win32';
break;
default:
return null;
}
let archKey;
switch (cpuArch) {
case 'x64':
case 'x86_64':
archKey = 'x64';
break;
case 'arm64':
case 'aarch64':
archKey = 'arm64';
break;
default:
return null;
}
const ext = os === 'win32' ? '.exe' : '';
return `agent-browser-${osKey}-${archKey}${ext}`;
}
function main() {
const binaryName = getBinaryName();
if (!binaryName) {
console.error(`Error: Unsupported platform: ${platform()}-${arch()}`);
process.exit(1);
}
const binaryPath = join(__dirname, binaryName);
if (!existsSync(binaryPath)) {
console.error(`Error: No binary found for ${platform()}-${arch()}`);
console.error(`Expected: ${binaryPath}`);
console.error('');
console.error('Run "npm run build:native" to build for your platform,');
console.error('or reinstall the package to trigger the postinstall download.');
process.exit(1);
}
// Ensure binary is executable (fixes EACCES on macOS/Linux when postinstall didn't run,
// e.g., when using bun which blocks lifecycle scripts by default)
if (platform() !== 'win32') {
try {
accessSync(binaryPath, constants.X_OK);
} catch {
// Binary exists but isn't executable - fix it
try {
chmodSync(binaryPath, 0o755);
} catch (chmodErr) {
console.error(`Error: Cannot make binary executable: ${chmodErr.message}`);
console.error('Try running: chmod +x ' + binaryPath);
process.exit(1);
}
}
}
// Spawn the native binary with inherited stdio
const child = spawn(binaryPath, process.argv.slice(2), {
stdio: 'inherit',
windowsHide: false,
});
child.on('error', (err) => {
console.error(`Error executing binary: ${err.message}`);
process.exit(1);
});
child.on('close', (code) => {
process.exit(code ?? 0);
});
}
main();
+2669 -31
View File
File diff suppressed because it is too large Load Diff
+37 -3
View File
@@ -1,13 +1,35 @@
[package]
name = "agent-browser"
version = "0.4.2"
name = "agent-browser-stealth"
version = "0.16.3-fork.2"
edition = "2021"
description = "Fast browser automation CLI for AI agents"
description = "Stealth browser automation CLI for AI agents with anti-bot evasions"
license = "Apache-2.0"
[[bin]]
name = "agent-browser"
path = "src/main.rs"
[[bin]]
name = "agent-browser-stealth"
path = "src/main_stealth.rs"
[dependencies]
serde = { version = "1.0", features = ["derive"] }
serde_json = "1.0"
dirs = "5.0"
base64 = "0.22"
getrandom = "0.2"
tokio = { version = "1", features = ["rt-multi-thread", "macros", "net", "io-util", "time", "sync", "signal"] }
tokio-tungstenite = { version = "0.24", features = ["rustls-tls-webpki-roots"] }
futures-util = "0.3"
url = "2"
uuid = { version = "1", features = ["v4"] }
image = "0.25"
reqwest = { version = "0.12", default-features = false, features = ["json", "rustls-tls-webpki-roots"] }
sha2 = "0.10"
aes-gcm = "0.10"
async-trait = "0.1"
similar = "2"
[target.'cfg(unix)'.dependencies]
libc = "0.2"
@@ -15,8 +37,20 @@ libc = "0.2"
[target.'cfg(windows)'.dependencies]
windows-sys = { version = "0.52", features = ["Win32_System_Threading", "Win32_Foundation"] }
[build-dependencies]
serde = { version = "1.0", features = ["derive"] }
serde_json = "1.0"
[profile.release]
opt-level = 3
lto = true
codegen-units = 1
strip = true
[profile.ci]
inherits = "release"
lto = "thin"
codegen-units = 16
[patch.crates-io]
zune-jpeg = { path = "vendor/zune-jpeg" }
+481
View File
@@ -0,0 +1,481 @@
use std::collections::HashSet;
use std::env;
use std::fs;
use std::path::Path;
fn main() {
let protocol_dir = Path::new("cdp-protocol");
let out_dir = env::var("OUT_DIR").unwrap();
let out_path = Path::new(&out_dir).join("cdp_generated.rs");
let browser_path = protocol_dir.join("browser_protocol.json");
let js_path = protocol_dir.join("js_protocol.json");
if !browser_path.exists() && !js_path.exists() {
fs::write(
&out_path,
"// No protocol JSON files found in cdp-protocol/\n",
)
.unwrap();
return;
}
let mut all_domains: Vec<Domain> = Vec::new();
for path in [&browser_path, &js_path] {
if !path.exists() {
continue;
}
println!("cargo:rerun-if-changed={}", path.display());
let content = fs::read_to_string(path).unwrap();
let protocol: ProtocolSpec = match serde_json::from_str(&content) {
Ok(p) => p,
Err(e) => {
eprintln!("cargo:warning=Failed to parse {}: {}", path.display(), e);
continue;
}
};
all_domains.extend(protocol.domains);
}
// Collect all known type IDs per domain for cross-domain resolution
let mut domain_types: std::collections::HashMap<String, HashSet<String>> =
std::collections::HashMap::new();
for domain in &all_domains {
let mut types = HashSet::new();
for td in &domain.types {
types.insert(td.id.clone());
}
domain_types.insert(domain.domain.clone(), types);
}
// Known recursive struct fields that need Box wrapping
let recursive_fields: HashSet<(&str, &str, &str)> = [
("DOM", "Node", "contentDocument"),
("DOM", "Node", "templateContent"),
("DOM", "Node", "importedDocument"),
("Accessibility", "AXNode", "sources"),
("Runtime", "StackTrace", "parent"),
]
.into_iter()
.collect();
let mut output = String::new();
output.push_str("use serde::{Deserialize, Serialize};\n\n");
for domain in &all_domains {
generate_domain(domain, &domain_types, &recursive_fields, &mut output);
}
fs::write(&out_path, &output).unwrap();
}
#[allow(dead_code)]
#[derive(serde::Deserialize)]
struct ProtocolSpec {
domains: Vec<Domain>,
}
#[allow(dead_code)]
#[derive(serde::Deserialize, Clone)]
struct Domain {
domain: String,
#[serde(default)]
types: Vec<TypeDef>,
#[serde(default)]
commands: Vec<Command>,
#[serde(default)]
events: Vec<Event>,
}
#[allow(dead_code)]
#[derive(serde::Deserialize, Clone)]
struct TypeDef {
id: String,
#[serde(rename = "type", default)]
type_kind: String,
#[serde(default)]
properties: Vec<Property>,
#[serde(rename = "enum", default)]
enum_values: Vec<String>,
#[serde(default)]
description: Option<String>,
}
#[allow(dead_code)]
#[derive(serde::Deserialize, Clone)]
struct Command {
name: String,
#[serde(default)]
parameters: Vec<Property>,
#[serde(default)]
returns: Vec<Property>,
#[serde(default)]
description: Option<String>,
}
#[allow(dead_code)]
#[derive(serde::Deserialize, Clone)]
struct Event {
name: String,
#[serde(default)]
parameters: Vec<Property>,
#[serde(default)]
description: Option<String>,
}
#[allow(dead_code)]
#[derive(serde::Deserialize, Clone)]
struct Property {
name: String,
#[serde(rename = "type", default)]
type_kind: Option<String>,
#[serde(rename = "$ref", default)]
ref_type: Option<String>,
#[serde(default)]
optional: bool,
#[serde(default)]
description: Option<String>,
#[serde(default)]
items: Option<Box<ItemType>>,
#[serde(rename = "enum", default)]
enum_values: Vec<String>,
}
#[allow(dead_code)]
#[derive(serde::Deserialize, Clone)]
struct ItemType {
#[serde(rename = "type", default)]
type_kind: Option<String>,
#[serde(rename = "$ref", default)]
ref_type: Option<String>,
}
fn to_pascal_case(s: &str) -> String {
let mut result = String::new();
let mut capitalize = true;
for c in s.chars() {
if c == '_' || c == '-' || c == '.' {
capitalize = true;
} else if capitalize {
result.push(c.to_ascii_uppercase());
capitalize = false;
} else {
result.push(c);
}
}
result
}
fn to_snake_case(s: &str) -> String {
let mut result = String::new();
let chars: Vec<char> = s.chars().collect();
for (i, &c) in chars.iter().enumerate() {
if c.is_uppercase() && i > 0 {
// Only insert underscore at transitions from lowercase to uppercase,
// or when an uppercase sequence ends (e.g. "DOM" -> "dom", not "d_o_m")
let prev_upper = chars[i - 1].is_uppercase();
let next_lower = chars.get(i + 1).map_or(false, |n| n.is_lowercase());
if !prev_upper || next_lower {
result.push('_');
}
}
result.push(c.to_ascii_lowercase());
}
result
}
/// Resolve a $ref type reference. Cross-domain refs like "Page.FrameId" become
/// `super::cdp_page::FrameId`. Same-domain refs are used directly.
fn resolve_ref(
r: &str,
current_domain: &str,
domain_types: &std::collections::HashMap<String, HashSet<String>>,
) -> String {
let parts: Vec<&str> = r.split('.').collect();
if parts.len() == 2 {
let ref_domain = parts[0];
let ref_type = parts[1];
if ref_domain == current_domain {
to_pascal_case(ref_type)
} else {
// Check if this type actually exists in the referenced domain
if domain_types
.get(ref_domain)
.map_or(false, |t| t.contains(ref_type))
{
format!(
"super::cdp_{}::{}",
to_snake_case(ref_domain),
to_pascal_case(ref_type)
)
} else {
// Fall back to serde_json::Value for unknown cross-domain refs
"serde_json::Value".to_string()
}
}
} else {
to_pascal_case(r)
}
}
fn map_type_in_domain(
prop: &Property,
current_domain: &str,
domain_types: &std::collections::HashMap<String, HashSet<String>>,
) -> String {
if let Some(ref r) = prop.ref_type {
let type_name = resolve_ref(r, current_domain, domain_types);
if prop.optional {
format!("Option<{}>", type_name)
} else {
type_name
}
} else if let Some(ref t) = prop.type_kind {
let base = match t.as_str() {
"string" => "String".to_string(),
"integer" => "i64".to_string(),
"number" => "f64".to_string(),
"boolean" => "bool".to_string(),
"object" => "serde_json::Value".to_string(),
"any" => "serde_json::Value".to_string(),
"array" => {
if let Some(ref items) = prop.items {
let inner = if let Some(ref r) = items.ref_type {
resolve_ref(r, current_domain, domain_types)
} else {
match items.type_kind.as_deref().unwrap_or("any") {
"string" => "String".to_string(),
"integer" => "i64".to_string(),
"number" => "f64".to_string(),
"boolean" => "bool".to_string(),
_ => "serde_json::Value".to_string(),
}
};
format!("Vec<{}>", inner)
} else {
"Vec<serde_json::Value>".to_string()
}
}
_ => "serde_json::Value".to_string(),
};
if prop.optional {
format!("Option<{}>", base)
} else {
base
}
} else if prop.optional {
"Option<serde_json::Value>".to_string()
} else {
"serde_json::Value".to_string()
}
}
fn is_rust_keyword(s: &str) -> bool {
matches!(
s,
"type"
| "self"
| "Self"
| "super"
| "move"
| "ref"
| "fn"
| "mod"
| "use"
| "pub"
| "let"
| "mut"
| "const"
| "static"
| "if"
| "else"
| "for"
| "while"
| "loop"
| "match"
| "return"
| "break"
| "continue"
| "as"
| "in"
| "impl"
| "trait"
| "struct"
| "enum"
| "where"
| "async"
| "await"
| "dyn"
| "box"
| "yield"
| "override"
| "crate"
| "extern"
)
}
fn generate_domain(
domain: &Domain,
domain_types: &std::collections::HashMap<String, HashSet<String>>,
recursive_fields: &HashSet<(&str, &str, &str)>,
output: &mut String,
) {
let mod_name = to_snake_case(&domain.domain);
output.push_str(&format!(
"#[allow(dead_code, non_snake_case, non_camel_case_types, clippy::enum_variant_names)]\npub mod cdp_{} {{\n",
mod_name
));
output.push_str(" use super::*;\n\n");
for type_def in &domain.types {
if !type_def.enum_values.is_empty() {
// Deduplicate enum variants (some CDP enums have duplicated PascalCase forms)
let mut seen_variants = HashSet::new();
output.push_str(" #[derive(Debug, Clone, Serialize, Deserialize)]\n");
output.push_str(&format!(" pub enum {} {{\n", type_def.id));
for val in &type_def.enum_values {
let mut variant = to_pascal_case(val);
if variant == "Self" {
variant = "SelfValue".to_string();
}
if variant.chars().next().map_or(false, |c| c.is_ascii_digit()) {
variant = format!("V{}", variant);
}
if seen_variants.insert(variant.clone()) {
output.push_str(&format!(
" #[serde(rename = \"{}\")]\n {},\n",
val, variant
));
}
}
output.push_str(" }\n\n");
} else if type_def.type_kind == "object" && !type_def.properties.is_empty() {
output.push_str(
" #[derive(Debug, Clone, Serialize, Deserialize)]\n #[serde(rename_all = \"camelCase\")]\n",
);
output.push_str(&format!(" pub struct {} {{\n", type_def.id));
for prop in &type_def.properties {
let field_name = to_snake_case(&prop.name);
let field_name = if is_rust_keyword(&field_name) {
format!("r#{}", field_name)
} else {
field_name
};
let mut rust_type = map_type_in_domain(prop, &domain.domain, domain_types);
// Wrap recursive fields in Box
if recursive_fields.contains(&(
domain.domain.as_str(),
type_def.id.as_str(),
prop.name.as_str(),
)) {
if rust_type.starts_with("Option<") {
let inner = &rust_type[7..rust_type.len() - 1];
rust_type = format!("Option<Box<{}>>", inner);
} else {
rust_type = format!("Box<{}>", rust_type);
}
}
if prop.optional {
output
.push_str(" #[serde(skip_serializing_if = \"Option::is_none\")]\n");
}
output.push_str(&format!(" pub {}: {},\n", field_name, rust_type));
}
output.push_str(" }\n\n");
} else if type_def.type_kind == "object" && type_def.properties.is_empty() {
output.push_str(&format!(
" pub type {} = serde_json::Value;\n\n",
type_def.id
));
} else if type_def.type_kind == "array" {
output.push_str(&format!(
" pub type {} = Vec<serde_json::Value>;\n\n",
type_def.id
));
} else if type_def.type_kind == "string" && type_def.enum_values.is_empty() {
output.push_str(&format!(" pub type {} = String;\n\n", type_def.id));
} else if type_def.type_kind == "integer" {
output.push_str(&format!(" pub type {} = i64;\n\n", type_def.id));
} else if type_def.type_kind == "number" {
output.push_str(&format!(" pub type {} = f64;\n\n", type_def.id));
}
}
for cmd in &domain.commands {
let pascal_name = to_pascal_case(&cmd.name);
if !cmd.parameters.is_empty() {
output.push_str(
" #[derive(Debug, Clone, Serialize, Deserialize)]\n #[serde(rename_all = \"camelCase\")]\n",
);
output.push_str(&format!(" pub struct {}Params {{\n", pascal_name));
for param in &cmd.parameters {
let field_name = to_snake_case(&param.name);
let field_name = if is_rust_keyword(&field_name) {
format!("r#{}", field_name)
} else {
field_name
};
let rust_type = map_type_in_domain(param, &domain.domain, domain_types);
if param.optional {
output
.push_str(" #[serde(skip_serializing_if = \"Option::is_none\")]\n");
}
output.push_str(&format!(" pub {}: {},\n", field_name, rust_type));
}
output.push_str(" }\n\n");
}
if !cmd.returns.is_empty() {
output.push_str(
" #[derive(Debug, Clone, Serialize, Deserialize)]\n #[serde(rename_all = \"camelCase\")]\n",
);
output.push_str(&format!(" pub struct {}Result {{\n", pascal_name));
for ret in &cmd.returns {
let field_name = to_snake_case(&ret.name);
let field_name = if is_rust_keyword(&field_name) {
format!("r#{}", field_name)
} else {
field_name
};
let rust_type = map_type_in_domain(ret, &domain.domain, domain_types);
if ret.optional {
output
.push_str(" #[serde(skip_serializing_if = \"Option::is_none\")]\n");
}
output.push_str(&format!(" pub {}: {},\n", field_name, rust_type));
}
output.push_str(" }\n\n");
}
}
for event in &domain.events {
if !event.parameters.is_empty() {
let pascal_name = to_pascal_case(&event.name);
output.push_str(
" #[derive(Debug, Clone, Serialize, Deserialize)]\n #[serde(rename_all = \"camelCase\")]\n",
);
output.push_str(&format!(" pub struct {}Event {{\n", pascal_name));
for param in &event.parameters {
let field_name = to_snake_case(&param.name);
let field_name = if is_rust_keyword(&field_name) {
format!("r#{}", field_name)
} else {
field_name
};
let rust_type = map_type_in_domain(param, &domain.domain, domain_types);
if param.optional {
output
.push_str(" #[serde(skip_serializing_if = \"Option::is_none\")]\n");
}
output.push_str(&format!(" pub {}: {},\n", field_name, rust_type));
}
output.push_str(" }\n\n");
}
}
output.push_str("}\n\n");
}
File diff suppressed because it is too large Load Diff
File diff suppressed because it is too large Load Diff
+158
View File
@@ -0,0 +1,158 @@
//! Color output utilities respecting NO_COLOR environment variable.
//!
//! When the NO_COLOR environment variable is present (regardless of value),
//! all color formatting is disabled per https://no-color.org/
use std::env;
use std::sync::OnceLock;
/// Returns true if color output is enabled (NO_COLOR is NOT set)
pub fn is_enabled() -> bool {
static COLORS_ENABLED: OnceLock<bool> = OnceLock::new();
*COLORS_ENABLED.get_or_init(|| env::var("NO_COLOR").is_err())
}
/// Format text in red (errors)
pub fn red(text: &str) -> String {
if is_enabled() {
format!("\x1b[31m{}\x1b[0m", text)
} else {
text.to_string()
}
}
/// Format text in green (success)
pub fn green(text: &str) -> String {
if is_enabled() {
format!("\x1b[32m{}\x1b[0m", text)
} else {
text.to_string()
}
}
/// Format text in yellow (warnings)
pub fn yellow(text: &str) -> String {
if is_enabled() {
format!("\x1b[33m{}\x1b[0m", text)
} else {
text.to_string()
}
}
/// Format text in cyan (info/progress)
pub fn cyan(text: &str) -> String {
if is_enabled() {
format!("\x1b[36m{}\x1b[0m", text)
} else {
text.to_string()
}
}
/// Format text in bold
pub fn bold(text: &str) -> String {
if is_enabled() {
format!("\x1b[1m{}\x1b[0m", text)
} else {
text.to_string()
}
}
/// Format text in dim
pub fn dim(text: &str) -> String {
if is_enabled() {
format!("\x1b[2m{}\x1b[0m", text)
} else {
text.to_string()
}
}
/// Red X error indicator
pub fn error_indicator() -> &'static str {
static INDICATOR: OnceLock<String> = OnceLock::new();
INDICATOR.get_or_init(|| {
if is_enabled() {
"\x1b[31m✗\x1b[0m".to_string()
} else {
"".to_string()
}
})
}
/// Green checkmark success indicator
pub fn success_indicator() -> &'static str {
static INDICATOR: OnceLock<String> = OnceLock::new();
INDICATOR.get_or_init(|| {
if is_enabled() {
"\x1b[32m✓\x1b[0m".to_string()
} else {
"".to_string()
}
})
}
/// Yellow warning indicator
pub fn warning_indicator() -> &'static str {
static INDICATOR: OnceLock<String> = OnceLock::new();
INDICATOR.get_or_init(|| {
if is_enabled() {
"\x1b[33m⚠\x1b[0m".to_string()
} else {
"".to_string()
}
})
}
/// Get console log color prefix by level
pub fn console_level_prefix(level: &str) -> String {
if !is_enabled() {
return format!("[{}]", level);
}
let color = match level {
"error" => "\x1b[31m",
"warning" => "\x1b[33m",
"info" => "\x1b[36m",
_ => "",
};
if color.is_empty() {
format!("[{}]", level)
} else {
format!("{}[{}]\x1b[0m", color, level)
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_red_contains_ansi_codes() {
// Test the format structure (actual color depends on NO_COLOR env)
let formatted = format!("\x1b[31m{}\x1b[0m", "error");
assert!(formatted.contains("\x1b[31m"));
assert!(formatted.contains("\x1b[0m"));
}
#[test]
fn test_green_contains_ansi_codes() {
let formatted = format!("\x1b[32m{}\x1b[0m", "success");
assert!(formatted.contains("\x1b[32m"));
}
#[test]
fn test_console_level_prefix_contains_level() {
// Regardless of color state, the level text should be present
assert!(console_level_prefix("error").contains("error"));
assert!(console_level_prefix("warning").contains("warning"));
assert!(console_level_prefix("info").contains("info"));
assert!(console_level_prefix("log").contains("log"));
}
#[test]
fn test_indicators_contain_symbols() {
// Regardless of color state, symbols should be present
assert!(error_indicator().contains('✗'));
assert!(success_indicator().contains('✓'));
assert!(warning_indicator().contains('⚠'));
}
}
+2822 -151
View File
File diff suppressed because it is too large Load Diff
+535 -60
View File
@@ -81,21 +81,68 @@ impl Connection {
}
}
/// Get the base directory for socket/pid files.
/// Priority: AGENT_BROWSER_SOCKET_DIR > XDG_RUNTIME_DIR > ~/.agent-browser > tmpdir
pub fn get_socket_dir() -> PathBuf {
// 1. Explicit override (ignore empty string)
if let Ok(dir) = env::var("AGENT_BROWSER_SOCKET_DIR") {
if !dir.is_empty() {
return PathBuf::from(dir);
}
}
// 2. XDG_RUNTIME_DIR (Linux standard, ignore empty string)
if let Ok(runtime_dir) = env::var("XDG_RUNTIME_DIR") {
if !runtime_dir.is_empty() {
return PathBuf::from(runtime_dir).join("agent-browser");
}
}
// 3. Home directory fallback (like Docker Desktop's ~/.docker/run/)
if let Some(home) = dirs::home_dir() {
return home.join(".agent-browser");
}
// 4. Last resort: temp dir
env::temp_dir().join("agent-browser")
}
#[cfg(unix)]
fn get_socket_path(session: &str) -> PathBuf {
let tmp = env::temp_dir();
tmp.join(format!("agent-browser-{}.sock", session))
get_socket_dir().join(format!("{}.sock", session))
}
fn get_pid_path(session: &str) -> PathBuf {
let tmp = env::temp_dir();
tmp.join(format!("agent-browser-{}.pid", session))
get_socket_dir().join(format!("{}.pid", session))
}
/// Clean up stale socket and PID files for a session
fn cleanup_stale_files(session: &str) {
// Never delete files for a live daemon. A missing PID file can happen in
// race scenarios, but the socket is authoritative for liveness.
if daemon_ready(session) {
return;
}
let pid_path = get_pid_path(session);
let _ = fs::remove_file(&pid_path);
#[cfg(unix)]
{
let socket_path = get_socket_path(session);
let _ = fs::remove_file(&socket_path);
}
#[cfg(windows)]
{
let port_path = get_port_path(session);
let _ = fs::remove_file(&port_path);
}
}
#[cfg(windows)]
fn get_port_path(session: &str) -> PathBuf {
let tmp = env::temp_dir();
tmp.join(format!("agent-browser-{}.port", session))
get_socket_dir().join(format!("{}.port", session))
}
#[cfg(windows)]
@@ -104,43 +151,16 @@ fn get_port_for_session(session: &str) -> u16 {
for c in session.chars() {
hash = ((hash << 5).wrapping_sub(hash)).wrapping_add(c as i32);
}
49152 + ((hash.abs() as u16) % 16383)
}
#[cfg(unix)]
fn is_daemon_running(session: &str) -> bool {
let pid_path = get_pid_path(session);
if !pid_path.exists() {
return false;
}
if let Ok(pid_str) = fs::read_to_string(&pid_path) {
if let Ok(pid) = pid_str.trim().parse::<i32>() {
unsafe {
return libc::kill(pid, 0) == 0;
}
}
}
false
}
#[cfg(windows)]
fn is_daemon_running(session: &str) -> bool {
let pid_path = get_pid_path(session);
if !pid_path.exists() {
return false;
}
let port = get_port_for_session(session);
TcpStream::connect_timeout(
&format!("127.0.0.1:{}", port).parse().unwrap(),
Duration::from_millis(100),
)
.is_ok()
// Correct logic: first take absolute modulo, then cast to u16
// Using unsigned_abs() to safely handle i32::MIN
49152 + ((hash.unsigned_abs() as u32 % 16383) as u16)
}
fn daemon_ready(session: &str) -> bool {
#[cfg(unix)]
{
get_socket_path(session).exists()
let socket_path = get_socket_path(session);
UnixStream::connect(&socket_path).is_ok()
}
#[cfg(windows)]
{
@@ -153,30 +173,120 @@ fn daemon_ready(session: &str) -> bool {
}
}
pub fn ensure_daemon(session: &str, headed: bool) -> Result<(), String> {
if is_daemon_running(session) && daemon_ready(session) {
return Ok(());
/// Result of ensure_daemon indicating whether a new daemon was started
pub struct DaemonResult {
/// True if we connected to an existing daemon, false if we started a new one
pub already_running: bool,
}
#[allow(clippy::too_many_arguments)]
pub fn ensure_daemon(
session: &str,
headed: bool,
executable_path: Option<&str>,
extensions: &[String],
args: Option<&str>,
user_agent: Option<&str>,
proxy: Option<&str>,
proxy_bypass: Option<&str>,
ignore_https_errors: bool,
allow_file_access: bool,
state: Option<&str>,
provider: Option<&str>,
device: Option<&str>,
session_name: Option<&str>,
debug: bool,
download_path: Option<&str>,
tab_group: Option<&str>,
tab_group_plugin_id: Option<&str>,
) -> Result<DaemonResult, String> {
// Socket readiness is the source of truth for a usable daemon.
// PID files can be missing/stale under concurrent start/stop races.
if daemon_ready(session) {
// Double-check it's actually responsive by waiting and checking again
// This handles the race condition where daemon is shutting down
// (daemon has a 100ms shutdown delay, so we wait longer)
thread::sleep(Duration::from_millis(150));
if daemon_ready(session) {
return Ok(DaemonResult {
already_running: true,
});
}
}
// Clean up any stale socket/pid files before starting fresh
cleanup_stale_files(session);
// Ensure socket directory exists
let socket_dir = get_socket_dir();
if !socket_dir.exists() {
fs::create_dir_all(&socket_dir)
.map_err(|e| format!("Failed to create socket directory: {}", e))?;
}
// Pre-flight check: Validate socket path length (Unix limit is 104 bytes including null terminator)
#[cfg(unix)]
{
let socket_path = get_socket_path(session);
let path_len = socket_path.as_os_str().len();
if path_len > 103 {
return Err(format!(
"Session name '{}' is too long. Socket path would be {} bytes (max 103).\n\
Use a shorter session name or set AGENT_BROWSER_SOCKET_DIR to a shorter path.",
session, path_len
));
}
}
// Pre-flight check: Verify socket directory is writable
{
let test_file = socket_dir.join(".write_test");
match fs::write(&test_file, b"") {
Ok(_) => {
let _ = fs::remove_file(&test_file);
}
Err(e) => {
return Err(format!(
"Socket directory '{}' is not writable: {}",
socket_dir.display(),
e
));
}
}
}
let exe_path = env::current_exe().map_err(|e| e.to_string())?;
// Canonicalize to resolve symlinks (e.g., npm global bin symlink -> actual binary)
let exe_path = exe_path.canonicalize().unwrap_or(exe_path);
let exe_dir = exe_path.parent().unwrap();
let daemon_paths = [
let mut daemon_paths = vec![
exe_dir.join("daemon.js"),
exe_dir.join("../dist/daemon.js"),
PathBuf::from("dist/daemon.js"),
];
// Check AGENT_BROWSER_HOME environment variable
if let Ok(home) = env::var("AGENT_BROWSER_HOME") {
let home_path = PathBuf::from(&home);
daemon_paths.insert(0, home_path.join("dist/daemon.js"));
daemon_paths.insert(1, home_path.join("daemon.js"));
}
let daemon_path = daemon_paths
.iter()
.find(|p| p.exists())
.ok_or("Daemon not found. Run from project directory or ensure daemon.js is alongside binary.")?;
.ok_or("Daemon not found. Set AGENT_BROWSER_HOME environment variable or run from project directory.")?;
// Keep handle to detect early daemon exit and surface startup errors.
#[allow(unused_assignments)]
let mut daemon_child: Option<std::process::Child> = None;
// Spawn daemon as a fully detached background process
#[cfg(unix)]
{
use std::os::unix::process::CommandExt;
let mut cmd = Command::new("node");
cmd.arg(daemon_path)
.env("AGENT_BROWSER_DAEMON", "1")
@@ -186,6 +296,68 @@ pub fn ensure_daemon(session: &str, headed: bool) -> Result<(), String> {
cmd.env("AGENT_BROWSER_HEADED", "1");
}
if let Some(path) = executable_path {
cmd.env("AGENT_BROWSER_EXECUTABLE_PATH", path);
}
if !extensions.is_empty() {
cmd.env("AGENT_BROWSER_EXTENSIONS", extensions.join(","));
}
if let Some(a) = args {
cmd.env("AGENT_BROWSER_ARGS", a);
}
if let Some(ua) = user_agent {
cmd.env("AGENT_BROWSER_USER_AGENT", ua);
}
if let Some(p) = proxy {
cmd.env("AGENT_BROWSER_PROXY", p);
}
if let Some(pb) = proxy_bypass {
cmd.env("AGENT_BROWSER_PROXY_BYPASS", pb);
}
if ignore_https_errors {
cmd.env("AGENT_BROWSER_IGNORE_HTTPS_ERRORS", "1");
}
if allow_file_access {
cmd.env("AGENT_BROWSER_ALLOW_FILE_ACCESS", "1");
}
if let Some(st) = state {
cmd.env("AGENT_BROWSER_STATE", st);
}
if let Some(p) = provider {
cmd.env("AGENT_BROWSER_PROVIDER", p);
}
if let Some(d) = device {
cmd.env("AGENT_BROWSER_IOS_DEVICE", d);
}
if let Some(sn) = session_name {
cmd.env("AGENT_BROWSER_SESSION_NAME", sn);
}
cmd.env("AGENT_BROWSER_STEALTH", "1");
if debug {
cmd.env("AGENT_BROWSER_DEBUG", "1");
}
if let Some(dp) = download_path {
cmd.env("AGENT_BROWSER_DOWNLOAD_PATH", dp);
}
if let Some(tg) = tab_group {
cmd.env("AGENT_BROWSER_TAB_GROUP", tg);
}
if let Some(plugin_id) = tab_group_plugin_id {
cmd.env("AGENT_BROWSER_TAB_GROUP_PLUGIN_ID", plugin_id);
}
// Create new process group and session to fully detach
unsafe {
cmd.pre_exec(|| {
@@ -195,17 +367,21 @@ pub fn ensure_daemon(session: &str, headed: bool) -> Result<(), String> {
});
}
cmd.stdin(Stdio::null())
.stdout(Stdio::null())
.stderr(Stdio::null())
.spawn()
.map_err(|e| format!("Failed to start daemon: {}", e))?;
daemon_child = Some(
cmd.stdin(Stdio::null())
.stdout(Stdio::null())
.stderr(Stdio::piped())
.spawn()
.map_err(|e| format!("Failed to start daemon: {}", e))?,
);
}
#[cfg(windows)]
{
use std::os::windows::process::CommandExt;
// On Windows, call node directly. Command::new handles PATH resolution (node.exe or node.cmd)
// and automatically quotes arguments containing spaces.
let mut cmd = Command::new("node");
cmd.arg(daemon_path)
.env("AGENT_BROWSER_DAEMON", "1")
@@ -215,26 +391,111 @@ pub fn ensure_daemon(session: &str, headed: bool) -> Result<(), String> {
cmd.env("AGENT_BROWSER_HEADED", "1");
}
if let Some(path) = executable_path {
cmd.env("AGENT_BROWSER_EXECUTABLE_PATH", path);
}
if !extensions.is_empty() {
cmd.env("AGENT_BROWSER_EXTENSIONS", extensions.join(","));
}
if let Some(a) = args {
cmd.env("AGENT_BROWSER_ARGS", a);
}
if let Some(ua) = user_agent {
cmd.env("AGENT_BROWSER_USER_AGENT", ua);
}
if let Some(p) = proxy {
cmd.env("AGENT_BROWSER_PROXY", p);
}
if let Some(pb) = proxy_bypass {
cmd.env("AGENT_BROWSER_PROXY_BYPASS", pb);
}
if ignore_https_errors {
cmd.env("AGENT_BROWSER_IGNORE_HTTPS_ERRORS", "1");
}
if allow_file_access {
cmd.env("AGENT_BROWSER_ALLOW_FILE_ACCESS", "1");
}
if let Some(st) = state {
cmd.env("AGENT_BROWSER_STATE", st);
}
if let Some(p) = provider {
cmd.env("AGENT_BROWSER_PROVIDER", p);
}
if let Some(d) = device {
cmd.env("AGENT_BROWSER_IOS_DEVICE", d);
}
if let Some(sn) = session_name {
cmd.env("AGENT_BROWSER_SESSION_NAME", sn);
}
cmd.env("AGENT_BROWSER_STEALTH", "1");
if debug {
cmd.env("AGENT_BROWSER_DEBUG", "1");
}
if let Some(dp) = download_path {
cmd.env("AGENT_BROWSER_DOWNLOAD_PATH", dp);
}
if let Some(tg) = tab_group {
cmd.env("AGENT_BROWSER_TAB_GROUP", tg);
}
if let Some(plugin_id) = tab_group_plugin_id {
cmd.env("AGENT_BROWSER_TAB_GROUP_PLUGIN_ID", plugin_id);
}
// CREATE_NEW_PROCESS_GROUP | DETACHED_PROCESS
const CREATE_NEW_PROCESS_GROUP: u32 = 0x00000200;
const DETACHED_PROCESS: u32 = 0x00000008;
cmd.creation_flags(CREATE_NEW_PROCESS_GROUP | DETACHED_PROCESS)
.stdin(Stdio::null())
.stdout(Stdio::null())
.stderr(Stdio::null())
.spawn()
.map_err(|e| format!("Failed to start daemon: {}", e))?;
daemon_child = Some(
cmd.creation_flags(CREATE_NEW_PROCESS_GROUP | DETACHED_PROCESS)
.stdin(Stdio::null())
.stdout(Stdio::null())
.stderr(Stdio::piped())
.spawn()
.map_err(|e| format!("Failed to start daemon: {}", e))?,
);
}
for _ in 0..50 {
if daemon_ready(session) {
return Ok(());
return Ok(DaemonResult {
already_running: false,
});
}
// Surface daemon startup stderr instead of returning an opaque timeout.
if let Some(ref mut child) = daemon_child {
if let Ok(Some(_)) = child.try_wait() {
let mut stderr_output = String::new();
if let Some(mut stderr) = child.stderr.take() {
let _ = stderr.read_to_string(&mut stderr_output);
}
let stderr_trimmed = stderr_output.trim();
if !stderr_trimmed.is_empty() {
return Err(format!("Daemon failed to start: {}", stderr_trimmed));
}
return Err("Daemon failed to start: process exited during startup".to_string());
}
}
thread::sleep(Duration::from_millis(100));
}
Err("Daemon failed to start".to_string())
Err(format!(
"Daemon failed to start (socket: {})",
get_socket_dir().join(format!("{}.sock", session)).display()
))
}
fn connect(session: &str) -> Result<Connection, String> {
@@ -255,12 +516,65 @@ fn connect(session: &str) -> Result<Connection, String> {
}
pub fn send_command(cmd: Value, session: &str) -> Result<Response, String> {
// Retry logic for transient errors (EAGAIN/EWOULDBLOCK/connection issues)
const MAX_RETRIES: u32 = 5;
const RETRY_DELAY_MS: u64 = 200;
let mut last_error = String::new();
for attempt in 0..MAX_RETRIES {
if attempt > 0 {
thread::sleep(Duration::from_millis(RETRY_DELAY_MS * (attempt as u64)));
}
match send_command_once(&cmd, session) {
Ok(response) => return Ok(response),
Err(e) => {
if is_transient_error(&e) {
last_error = e;
continue;
}
// Non-transient error, fail immediately
return Err(e);
}
}
}
Err(format!(
"{} (after {} retries - daemon may be busy or unresponsive)",
last_error, MAX_RETRIES
))
}
/// Check if an error is transient and worth retrying.
/// Transient errors include:
/// - EAGAIN/EWOULDBLOCK (os error 35 on macOS, 11 on Linux)
/// - EOF errors (daemon closed connection before responding)
/// - Connection reset/broken pipe (daemon crashed or restarting)
/// - Connection refused/socket not found (daemon still starting)
fn is_transient_error(error: &str) -> bool {
error.contains("os error 35") // EAGAIN on macOS
|| error.contains("os error 11") // EAGAIN on Linux
|| error.contains("WouldBlock")
|| error.contains("Resource temporarily unavailable")
|| error.contains("EOF")
|| error.contains("line 1 column 0") // Empty JSON response
|| error.contains("Connection reset")
|| error.contains("Broken pipe")
|| error.contains("os error 54") // Connection reset by peer (macOS)
|| error.contains("os error 104") // Connection reset by peer (Linux)
|| error.contains("os error 2") // No such file or directory (socket gone)
|| error.contains("os error 61") // Connection refused (macOS)
|| error.contains("os error 111") // Connection refused (Linux)
}
fn send_command_once(cmd: &Value, session: &str) -> Result<Response, String> {
let mut stream = connect(session)?;
stream.set_read_timeout(Some(Duration::from_secs(30))).ok();
stream.set_write_timeout(Some(Duration::from_secs(5))).ok();
let mut json_str = serde_json::to_string(&cmd).map_err(|e| e.to_string())?;
let mut json_str = serde_json::to_string(cmd).map_err(|e| e.to_string())?;
json_str.push('\n');
stream
@@ -275,3 +589,164 @@ pub fn send_command(cmd: Value, session: &str) -> Result<Response, String> {
serde_json::from_str(&response_line).map_err(|e| format!("Invalid response: {}", e))
}
#[cfg(test)]
mod tests {
use super::*;
use crate::test_utils::EnvGuard;
#[test]
fn test_get_socket_dir_explicit_override() {
let _guard = EnvGuard::new(&["AGENT_BROWSER_SOCKET_DIR", "XDG_RUNTIME_DIR"]);
_guard.set("AGENT_BROWSER_SOCKET_DIR", "/custom/socket/path");
_guard.remove("XDG_RUNTIME_DIR");
assert_eq!(get_socket_dir(), PathBuf::from("/custom/socket/path"));
}
#[test]
fn test_get_socket_dir_ignores_empty_socket_dir() {
let _guard = EnvGuard::new(&["AGENT_BROWSER_SOCKET_DIR", "XDG_RUNTIME_DIR"]);
_guard.set("AGENT_BROWSER_SOCKET_DIR", "");
_guard.remove("XDG_RUNTIME_DIR");
assert!(get_socket_dir()
.to_string_lossy()
.ends_with(".agent-browser"));
}
#[test]
fn test_get_socket_dir_xdg_runtime() {
let _guard = EnvGuard::new(&["AGENT_BROWSER_SOCKET_DIR", "XDG_RUNTIME_DIR"]);
_guard.remove("AGENT_BROWSER_SOCKET_DIR");
_guard.set("XDG_RUNTIME_DIR", "/run/user/1000");
assert_eq!(
get_socket_dir(),
PathBuf::from("/run/user/1000/agent-browser")
);
}
#[test]
fn test_get_socket_dir_ignores_empty_xdg_runtime() {
let _guard = EnvGuard::new(&["AGENT_BROWSER_SOCKET_DIR", "XDG_RUNTIME_DIR"]);
_guard.set("AGENT_BROWSER_SOCKET_DIR", "");
_guard.set("XDG_RUNTIME_DIR", "");
assert!(get_socket_dir()
.to_string_lossy()
.ends_with(".agent-browser"));
}
#[test]
fn test_get_socket_dir_home_fallback() {
let _guard = EnvGuard::new(&["AGENT_BROWSER_SOCKET_DIR", "XDG_RUNTIME_DIR"]);
_guard.remove("AGENT_BROWSER_SOCKET_DIR");
_guard.remove("XDG_RUNTIME_DIR");
let result = get_socket_dir();
assert!(result.to_string_lossy().ends_with(".agent-browser"));
assert!(
result.to_string_lossy().contains("home") || result.to_string_lossy().contains("Users")
);
}
// === Transient Error Detection Tests ===
#[test]
fn test_is_transient_error_eagain_macos() {
assert!(is_transient_error(
"Failed to read: Resource temporarily unavailable (os error 35)"
));
}
#[test]
fn test_is_transient_error_eagain_linux() {
assert!(is_transient_error(
"Failed to read: Resource temporarily unavailable (os error 11)"
));
}
#[test]
fn test_is_transient_error_would_block() {
assert!(is_transient_error("operation WouldBlock"));
}
#[test]
fn test_is_transient_error_resource_unavailable() {
assert!(is_transient_error("Resource temporarily unavailable"));
}
#[test]
fn test_is_transient_error_eof() {
assert!(is_transient_error(
"Invalid response: EOF while parsing a value at line 1 column 0"
));
}
#[test]
fn test_is_transient_error_empty_json() {
assert!(is_transient_error(
"Invalid response: expected value at line 1 column 0"
));
}
#[test]
fn test_is_transient_error_connection_reset() {
assert!(is_transient_error("Connection reset by peer"));
}
#[test]
fn test_is_transient_error_broken_pipe() {
assert!(is_transient_error("Broken pipe"));
}
#[test]
fn test_is_transient_error_connection_reset_macos() {
assert!(is_transient_error(
"Failed to send: Connection reset by peer (os error 54)"
));
}
#[test]
fn test_is_transient_error_connection_reset_linux() {
assert!(is_transient_error(
"Failed to send: Connection reset by peer (os error 104)"
));
}
#[test]
fn test_is_transient_error_socket_not_found() {
assert!(is_transient_error(
"Failed to connect: No such file or directory (os error 2)"
));
}
#[test]
fn test_is_transient_error_connection_refused_macos() {
assert!(is_transient_error(
"Failed to connect: Connection refused (os error 61)"
));
}
#[test]
fn test_is_transient_error_connection_refused_linux() {
assert!(is_transient_error(
"Failed to connect: Connection refused (os error 111)"
));
}
#[test]
fn test_is_transient_error_non_transient() {
// These should NOT be considered transient
assert!(!is_transient_error("Unknown command: foo"));
assert!(!is_transient_error("Invalid JSON syntax"));
assert!(!is_transient_error("Permission denied"));
assert!(!is_transient_error("Daemon not found"));
}
}
+1290 -18
View File
File diff suppressed because it is too large Load Diff
+57 -13
View File
@@ -1,3 +1,4 @@
use crate::color;
use std::process::{exit, Command, Stdio};
pub fn run_install(with_deps: bool) {
@@ -5,9 +6,15 @@ pub fn run_install(with_deps: bool) {
if is_linux {
if with_deps {
println!("\x1b[36mInstalling system dependencies...\x1b[0m");
println!("{}", color::cyan("Installing system dependencies..."));
let (pkg_mgr, deps) = if which_exists("apt-get") {
let libasound = if package_exists_apt("libasound2t64") {
"libasound2t64"
} else {
"libasound2"
};
(
"apt-get",
vec![
@@ -30,7 +37,7 @@ pub fn run_install(with_deps: bool) {
"libcairo2",
"libgdk-pixbuf-2.0-0",
"libxrender1",
"libasound2",
libasound,
"libfreetype6",
"libfontconfig1",
"libdbus-1-3",
@@ -93,7 +100,10 @@ pub fn run_install(with_deps: bool) {
],
)
} else {
eprintln!("\x1b[31m✗\x1b[0m No supported package manager found (apt-get, dnf, or yum)");
eprintln!(
"{} No supported package manager found (apt-get, dnf, or yum)",
color::error_indicator()
);
exit(1);
};
@@ -112,45 +122,68 @@ pub fn run_install(with_deps: bool) {
match status {
Ok(s) if s.success() => {
println!("\x1b[32m✓\x1b[0m System dependencies installed")
println!("{} System dependencies installed", color::success_indicator())
}
Ok(_) => eprintln!(
"\x1b[33m⚠\x1b[0m Failed to install some dependencies. You may need to run manually with sudo."
"{} Failed to install some dependencies. You may need to run manually with sudo.",
color::warning_indicator()
),
Err(e) => eprintln!("\x1b[33m⚠\x1b[0m Could not run install command: {}", e),
Err(e) => eprintln!("{} Could not run install command: {}", color::warning_indicator(), e),
}
} else {
println!("\x1b[33m⚠\x1b[0m Linux detected. If browser fails to launch, run:");
println!(
"{} Linux detected. If browser fails to launch, run:",
color::warning_indicator()
);
println!(" agent-browser install --with-deps");
println!(" or: npx playwright install-deps chromium");
println!();
}
}
println!("\x1b[36mInstalling Chromium browser...\x1b[0m");
println!("{}", color::cyan("Installing Chromium browser..."));
// On Windows, we need to use cmd.exe to run npx because npx is actually npx.cmd
// and Command::new() doesn't resolve .cmd files the way the shell does.
// Pass the entire command as a single string to /c to handle paths with spaces.
#[cfg(windows)]
let status = Command::new("cmd")
.args(["/c", "npx playwright install chromium"])
.status();
#[cfg(not(windows))]
let status = Command::new("npx")
.args(["playwright", "install", "chromium"])
.status();
match status {
Ok(s) if s.success() => {
println!("\x1b[32m✓\x1b[0m Chromium installed successfully");
println!(
"{} Chromium installed successfully",
color::success_indicator()
);
if is_linux && !with_deps {
println!();
println!("\x1b[33mNote:\x1b[0m If you see \"shared library\" errors when running, use:");
println!(
"{} If you see \"shared library\" errors when running, use:",
color::yellow("Note:")
);
println!(" agent-browser install --with-deps");
}
}
Ok(_) => {
eprintln!("\x1b[31m✗\x1b[0m Failed to install browser");
eprintln!("{} Failed to install browser", color::error_indicator());
if is_linux {
println!("\x1b[33mTip:\x1b[0m Try installing system dependencies first:");
println!(
"{} Try installing system dependencies first:",
color::yellow("Tip:")
);
println!(" agent-browser install --with-deps");
}
exit(1);
}
Err(e) => {
eprintln!("\x1b[31m✗\x1b[0m Failed to run npx: {}", e);
eprintln!("{} Failed to run npx: {}", color::error_indicator(), e);
eprintln!("Make sure Node.js is installed and npx is in your PATH");
exit(1);
}
@@ -179,3 +212,14 @@ fn which_exists(cmd: &str) -> bool {
.unwrap_or(false)
}
}
fn package_exists_apt(pkg: &str) -> bool {
Command::new("apt-cache")
.arg("show")
.arg(pkg)
.stdout(Stdio::null())
.stderr(Stdio::null())
.status()
.map(|s| s.success())
.unwrap_or(false)
}
+720 -33
View File
@@ -1,55 +1,88 @@
mod color;
mod commands;
mod connection;
mod flags;
mod install;
mod output;
#[cfg(test)]
mod test_utils;
mod validation;
use serde_json::json;
use std::env;
use std::fs;
use std::process::exit;
#[cfg(unix)]
use libc;
#[cfg(windows)]
use windows_sys::Win32::Foundation::CloseHandle;
#[cfg(windows)]
use windows_sys::Win32::System::Threading::{OpenProcess, PROCESS_QUERY_LIMITED_INFORMATION};
use commands::{gen_id, parse_command, ParseError};
use connection::{ensure_daemon, send_command};
use connection::{ensure_daemon, get_socket_dir, send_command};
use flags::{clean_args, parse_flags};
use install::run_install;
use output::{print_help, print_response};
use output::{print_command_help, print_help, print_response, print_version};
fn parse_proxy(proxy_str: &str) -> serde_json::Value {
let Some(protocol_end) = proxy_str.find("://") else {
return json!({ "server": proxy_str });
};
let protocol = &proxy_str[..protocol_end + 3];
let rest = &proxy_str[protocol_end + 3..];
let Some(at_pos) = rest.rfind('@') else {
return json!({ "server": proxy_str });
};
let creds = &rest[..at_pos];
let server_part = &rest[at_pos + 1..];
let server = format!("{}{}", protocol, server_part);
let Some(colon_pos) = creds.find(':') else {
return json!({
"server": server,
"username": creds,
"password": ""
});
};
json!({
"server": server,
"username": &creds[..colon_pos],
"password": &creds[colon_pos + 1..]
})
}
fn run_session(args: &[String], session: &str, json_mode: bool) {
let subcommand = args.get(1).map(|s| s.as_str());
match subcommand {
Some("list") => {
let tmp = env::temp_dir();
let socket_dir = get_socket_dir();
let mut sessions: Vec<String> = Vec::new();
if let Ok(entries) = fs::read_dir(&tmp) {
if let Ok(entries) = fs::read_dir(&socket_dir) {
for entry in entries.flatten() {
let name = entry.file_name().to_string_lossy().to_string();
// Look for socket files (Unix) or pid files
if name.starts_with("agent-browser-") && name.ends_with(".pid") {
let session_name = name
.strip_prefix("agent-browser-")
.and_then(|s| s.strip_suffix(".pid"))
.unwrap_or("");
// Look for pid files in socket directory
if name.ends_with(".pid") {
let session_name = name.strip_suffix(".pid").unwrap_or("");
if !session_name.is_empty() {
// Check if session is actually running
let pid_path = tmp.join(&name);
let pid_path = socket_dir.join(&name);
if let Ok(pid_str) = fs::read_to_string(&pid_path) {
if let Ok(pid) = pid_str.trim().parse::<u32>() {
#[cfg(unix)]
let running = unsafe { libc::kill(pid as i32, 0) == 0 };
let running = unsafe {
libc::kill(pid as i32, 0) == 0
|| std::io::Error::last_os_error().raw_os_error()
!= Some(libc::ESRCH)
};
#[cfg(windows)]
let running = unsafe {
let handle = OpenProcess(PROCESS_QUERY_LIMITED_INFORMATION, 0, pid);
let handle =
OpenProcess(PROCESS_QUERY_LIMITED_INFORMATION, 0, pid);
if handle != 0 {
CloseHandle(handle);
true
@@ -77,7 +110,11 @@ fn run_session(args: &[String], session: &str, json_mode: bool) {
} else {
println!("Active sessions:");
for s in &sessions {
let marker = if s == session { "" } else { " " };
let marker = if s == session {
color::cyan("")
} else {
" ".to_string()
};
println!("{} {}", marker, s);
}
}
@@ -94,24 +131,104 @@ fn run_session(args: &[String], session: &str, json_mode: bool) {
}
fn main() {
// Ignore SIGPIPE to prevent panic when piping to head/tail
#[cfg(unix)]
unsafe {
libc::signal(libc::SIGPIPE, libc::SIG_DFL);
}
let args: Vec<String> = env::args().skip(1).collect();
let flags = parse_flags(&args);
let clean = clean_args(&args);
if clean.is_empty() || args.iter().any(|a| a == "--help" || a == "-h") {
let has_help = args.iter().any(|a| a == "--help" || a == "-h");
let has_version = args.iter().any(|a| a == "--version" || a == "-V");
if has_help {
if let Some(cmd) = clean.first() {
if print_command_help(cmd) {
return;
}
}
print_help();
return;
}
if has_version {
print_version();
return;
}
if let Some(ref risk_mode) = flags.risk_mode {
if !matches!(risk_mode.as_str(), "off" | "warn" | "block") {
let msg = format!(
"Invalid --risk-mode value: {} (expected off, warn, or block)",
risk_mode
);
if flags.json {
println!(r#"{{"success":false,"error":"{}"}}"#, msg);
} else {
eprintln!("{} {}", color::error_indicator(), msg);
}
exit(1);
}
}
if args.iter().any(|a| a == "--profile") {
let msg =
"Project policy: --profile is forbidden. Use your existing browser and --session-name for state persistence.";
if flags.json {
println!(r#"{{"success":false,"error":"{}"}}"#, msg);
} else {
eprintln!("{} {}", color::error_indicator(), msg);
}
exit(1);
}
if env::var("AGENT_BROWSER_PROFILE").is_ok() {
let msg =
"Project policy: AGENT_BROWSER_PROFILE is forbidden. Remove it and use --session-name.";
if flags.json {
println!(r#"{{"success":false,"error":"{}"}}"#, msg);
} else {
eprintln!("{} {}", color::error_indicator(), msg);
}
exit(1);
}
if args.iter().any(|a| a == "--channel") {
let msg = "Project policy: --channel is forbidden. Browser selection follows your existing browser session.";
if flags.json {
println!(r#"{{"success":false,"error":"{}"}}"#, msg);
} else {
eprintln!("{} {}", color::error_indicator(), msg);
}
exit(1);
}
if env::var("AGENT_BROWSER_CHANNEL").is_ok() {
let msg =
"Project policy: AGENT_BROWSER_CHANNEL is forbidden. Remove it and use your existing browser session.";
if flags.json {
println!(r#"{{"success":false,"error":"{}"}}"#, msg);
} else {
eprintln!("{} {}", color::error_indicator(), msg);
}
exit(1);
}
if clean.is_empty() {
print_help();
return;
}
// Handle install separately
if clean.get(0).map(|s| s.as_str()) == Some("install") {
if clean.first().map(|s| s.as_str()) == Some("install") {
let with_deps = args.iter().any(|a| a == "--with-deps" || a == "-d");
run_install(with_deps);
return;
}
// Handle session separately (doesn't need daemon)
if clean.get(0).map(|s| s.as_str()) == Some("session") {
if clean.first().map(|s| s.as_str()) == Some("session") {
run_session(&clean, &flags.session, flags.json);
return;
}
@@ -124,6 +241,8 @@ fn main() {
ParseError::UnknownCommand { .. } => "unknown_command",
ParseError::UnknownSubcommand { .. } => "unknown_subcommand",
ParseError::MissingArguments { .. } => "missing_arguments",
ParseError::InvalidValue { .. } => "invalid_value",
ParseError::InvalidSessionName { .. } => "invalid_session_name",
};
println!(
r#"{{"success":false,"error":"{}","type":"{}"}}"#,
@@ -131,35 +250,544 @@ fn main() {
error_type
);
} else {
eprintln!("\x1b[31m{}\x1b[0m", e.format());
eprintln!("{}", color::red(&e.format()));
}
exit(1);
}
};
if let Err(e) = ensure_daemon(&flags.session, flags.headed) {
// Validate session name before starting daemon
if let Some(ref name) = flags.session_name {
if !validation::is_valid_session_name(name) {
let msg = validation::session_name_error(name);
if flags.json {
println!(
r#"{{"success":false,"error":"{}","type":"invalid_session_name"}}"#,
msg.replace('"', "\\\"")
);
} else {
eprintln!("{} {}", color::error_indicator(), msg);
}
exit(1);
}
}
let daemon_result = match ensure_daemon(
&flags.session,
flags.headed,
flags.executable_path.as_deref(),
&flags.extensions,
flags.args.as_deref(),
flags.user_agent.as_deref(),
flags.proxy.as_deref(),
flags.proxy_bypass.as_deref(),
flags.ignore_https_errors,
flags.allow_file_access,
flags.state.as_deref(),
flags.provider.as_deref(),
flags.device.as_deref(),
flags.session_name.as_deref(),
flags.debug,
flags.download_path.as_deref(),
flags.tab_group.as_deref(),
flags.tab_group_plugin_id.as_deref(),
) {
Ok(result) => result,
Err(e) => {
if flags.json {
println!(r#"{{"success":false,"error":"{}"}}"#, e);
} else {
eprintln!("{} {}", color::error_indicator(), e);
}
exit(1);
}
};
// Warn if launch-time options were explicitly passed via CLI but daemon was already running
// Only warn about flags that were passed on the command line, not those set via environment
// variables (since the daemon already uses the env vars when it starts).
if daemon_result.already_running {
let ignored_flags: Vec<&str> = [
if flags.cli_executable_path {
Some("--executable-path")
} else {
None
},
if flags.cli_extensions {
Some("--extension")
} else {
None
},
if flags.cli_state {
Some("--state")
} else {
None
},
if flags.cli_args { Some("--args") } else { None },
if flags.cli_user_agent {
Some("--user-agent")
} else {
None
},
if flags.cli_proxy {
Some("--proxy")
} else {
None
},
if flags.cli_proxy_bypass {
Some("--proxy-bypass")
} else {
None
},
flags.ignore_https_errors.then_some("--ignore-https-errors"),
flags.cli_allow_file_access.then_some("--allow-file-access"),
flags.cli_download_path.then_some("--download-path"),
flags.cli_tab_group.then_some("--tab-group"),
flags
.cli_tab_group_plugin_id
.then_some("--tab-group-plugin-id"),
]
.into_iter()
.flatten()
.collect();
if !ignored_flags.is_empty() && !flags.json {
eprintln!(
"{} {} ignored: daemon already running. Use 'agent-browser close' first to restart with new options.",
color::warning_indicator(),
ignored_flags.join(", ")
);
}
}
// Validate mutually exclusive options
if flags.cdp.is_some() && flags.provider.is_some() {
let msg = "Cannot use --cdp and -p/--provider together";
if flags.json {
println!(r#"{{"success":false,"error":"{}"}}"#, e);
println!(r#"{{"success":false,"error":"{}"}}"#, msg);
} else {
eprintln!("\x1b[31m✗\x1b[0m {}", e);
eprintln!("{} {}", color::error_indicator(), msg);
}
exit(1);
}
// If --headed flag is set, send launch command first to switch to headed mode
if flags.headed {
let launch_cmd = json!({ "id": gen_id(), "action": "launch", "headless": false });
if let Err(e) = send_command(launch_cmd, &flags.session) {
if !flags.json {
eprintln!("\x1b[33m⚠\x1b[0m Could not switch to headed mode: {}", e);
if flags.auto_connect && flags.cdp.is_some() {
let msg = "Cannot use --auto-connect and --cdp together";
if flags.json {
println!(r#"{{"success":false,"error":"{}"}}"#, msg);
} else {
eprintln!("{} {}", color::error_indicator(), msg);
}
exit(1);
}
if flags.auto_connect && flags.provider.is_some() {
let msg = "Cannot use --auto-connect and -p/--provider together";
if flags.json {
println!(r#"{{"success":false,"error":"{}"}}"#, msg);
} else {
eprintln!("{} {}", color::error_indicator(), msg);
}
exit(1);
}
if flags.provider.is_some() && !flags.extensions.is_empty() {
let msg = "Cannot use --extension with -p/--provider (extensions require local browser)";
if flags.json {
println!(r#"{{"success":false,"error":"{}"}}"#, msg);
} else {
eprintln!("{} {}", color::error_indicator(), msg);
}
exit(1);
}
if flags.cdp.is_some() && !flags.extensions.is_empty() {
let msg = "Cannot use --extension with --cdp (extensions require local browser)";
if flags.json {
println!(r#"{{"success":false,"error":"{}"}}"#, msg);
} else {
eprintln!("{} {}", color::error_indicator(), msg);
}
exit(1);
}
let mut attached_to_existing_browser = false;
// Auto-connect to existing browser
if flags.auto_connect {
let mut launch_cmd = json!({
"id": gen_id(),
"action": "launch",
"autoConnect": true
});
if flags.ignore_https_errors {
launch_cmd["ignoreHTTPSErrors"] = json!(true);
}
if let Some(ref cs) = flags.color_scheme {
launch_cmd["colorScheme"] = json!(cs);
}
if let Some(ref dp) = flags.download_path {
launch_cmd["downloadPath"] = json!(dp);
}
if let Some(ref tg) = flags.tab_group {
launch_cmd["tabGroup"] = json!(tg);
}
if let Some(ref plugin_id) = flags.tab_group_plugin_id {
launch_cmd["tabGroupPluginId"] = json!(plugin_id);
}
let err = match send_command(launch_cmd, &flags.session) {
Ok(resp) if resp.success => None,
Ok(resp) => Some(
resp.error
.unwrap_or_else(|| "Auto-connect failed".to_string()),
),
Err(e) => Some(e.to_string()),
};
if let Some(msg) = err {
if flags.json {
println!(r#"{{"success":false,"error":"{}"}}"#, msg);
} else {
eprintln!("{} {}", color::error_indicator(), msg);
}
exit(1);
}
attached_to_existing_browser = true;
}
// Connect via CDP if --cdp flag is set
// Accepts either a port number (e.g., "9222") or a full URL (e.g., "ws://..." or "wss://...")
if let Some(ref cdp_value) = flags.cdp {
let mut launch_cmd = if cdp_value.starts_with("ws://")
|| cdp_value.starts_with("wss://")
|| cdp_value.starts_with("http://")
|| cdp_value.starts_with("https://")
{
// It's a URL - use cdpUrl field
json!({
"id": gen_id(),
"action": "launch",
"cdpUrl": cdp_value
})
} else {
// It's a port number - validate and use cdpPort field
let cdp_port: u16 = match cdp_value.parse::<u32>() {
Ok(0) => {
let msg = "Invalid CDP port: port must be greater than 0".to_string();
if flags.json {
println!(r#"{{"success":false,"error":"{}"}}"#, msg);
} else {
eprintln!("{} {}", color::error_indicator(), msg);
}
exit(1);
}
Ok(p) if p > 65535 => {
let msg = format!(
"Invalid CDP port: {} is out of range (valid range: 1-65535)",
p
);
if flags.json {
println!(r#"{{"success":false,"error":"{}"}}"#, msg);
} else {
eprintln!("{} {}", color::error_indicator(), msg);
}
exit(1);
}
Ok(p) => p as u16,
Err(_) => {
let msg = format!(
"Invalid CDP value: '{}' is not a valid port number or URL",
cdp_value
);
if flags.json {
println!(r#"{{"success":false,"error":"{}"}}"#, msg);
} else {
eprintln!("{} {}", color::error_indicator(), msg);
}
exit(1);
}
};
json!({
"id": gen_id(),
"action": "launch",
"cdpPort": cdp_port
})
};
if flags.ignore_https_errors {
launch_cmd["ignoreHTTPSErrors"] = json!(true);
}
if let Some(ref cs) = flags.color_scheme {
launch_cmd["colorScheme"] = json!(cs);
}
if let Some(ref dp) = flags.download_path {
launch_cmd["downloadPath"] = json!(dp);
}
if let Some(ref tg) = flags.tab_group {
launch_cmd["tabGroup"] = json!(tg);
}
if let Some(ref plugin_id) = flags.tab_group_plugin_id {
launch_cmd["tabGroupPluginId"] = json!(plugin_id);
}
let err = match send_command(launch_cmd, &flags.session) {
Ok(resp) if resp.success => None,
Ok(resp) => Some(
resp.error
.unwrap_or_else(|| "CDP connection failed".to_string()),
),
Err(e) => Some(e.to_string()),
};
if let Some(msg) = err {
if flags.json {
println!(r#"{{"success":false,"error":"{}"}}"#, msg);
} else {
eprintln!("{} {}", color::error_indicator(), msg);
}
exit(1);
}
attached_to_existing_browser = true;
}
// Launch with cloud provider if -p flag is set
if let Some(ref provider) = flags.provider {
let mut launch_cmd = json!({
"id": gen_id(),
"action": "launch",
"provider": provider
});
if let Some(ref cs) = flags.color_scheme {
launch_cmd["colorScheme"] = json!(cs);
}
if let Some(ref tg) = flags.tab_group {
launch_cmd["tabGroup"] = json!(tg);
}
if let Some(ref plugin_id) = flags.tab_group_plugin_id {
launch_cmd["tabGroupPluginId"] = json!(plugin_id);
}
match send_command(launch_cmd, &flags.session) {
Ok(resp) => {
if !resp.success {
let msg = resp
.error
.unwrap_or_else(|| "Provider connection failed".to_string());
if flags.json {
println!(r#"{{"success":false,"error":"{}"}}"#, msg);
} else {
eprintln!("{} {}", color::error_indicator(), msg);
}
exit(1);
}
}
Err(e) => {
if flags.json {
println!(r#"{{"success":false,"error":"{}"}}"#, e);
} else {
eprintln!("{} {}", color::error_indicator(), e);
}
exit(1);
}
}
}
match send_command(cmd, &flags.session) {
// Project policy: when no explicit connection mode is provided,
// commands should attach to an existing browser.
// Try CDP :9333 first, then fall back to auto-connect discovery.
let can_try_default_cdp = flags.cdp.is_none()
&& !flags.auto_connect
&& flags.provider.is_none()
&& flags.executable_path.is_none()
&& flags.state.is_none()
&& flags.proxy.is_none()
&& flags.args.is_none()
&& flags.user_agent.is_none()
&& !flags.ignore_https_errors
&& !flags.allow_file_access
&& flags.extensions.is_empty();
if can_try_default_cdp {
let mut launch_cmd = json!({
"id": gen_id(),
"action": "launch",
"cdpPort": 9333
});
if let Some(ref cs) = flags.color_scheme {
launch_cmd["colorScheme"] = json!(cs);
}
if let Some(ref tg) = flags.tab_group {
launch_cmd["tabGroup"] = json!(tg);
}
if let Some(ref plugin_id) = flags.tab_group_plugin_id {
launch_cmd["tabGroupPluginId"] = json!(plugin_id);
}
if let Ok(resp) = send_command(launch_cmd, &flags.session) {
attached_to_existing_browser = resp.success;
}
if !attached_to_existing_browser {
let mut auto_connect_cmd = json!({
"id": gen_id(),
"action": "launch",
"autoConnect": true
});
if let Some(ref cs) = flags.color_scheme {
auto_connect_cmd["colorScheme"] = json!(cs);
}
if let Some(ref tg) = flags.tab_group {
auto_connect_cmd["tabGroup"] = json!(tg);
}
if let Some(ref plugin_id) = flags.tab_group_plugin_id {
auto_connect_cmd["tabGroupPluginId"] = json!(plugin_id);
}
if let Ok(resp) = send_command(auto_connect_cmd, &flags.session) {
attached_to_existing_browser = resp.success;
}
}
}
if can_try_default_cdp && !attached_to_existing_browser {
let msg = "Project policy requires using your existing browser. Could not connect to CDP at localhost:9333 and auto-discovery also failed. Start Chrome with remote debugging (for example, --remote-debugging-port=9333), or pass --cdp <port|url>.";
if flags.json {
println!(r#"{{"success":false,"error":"{}"}}"#, msg);
} else {
eprintln!("{} {}", color::error_indicator(), msg);
}
exit(1);
}
// Launch headed browser or configure browser options (without CDP or provider)
if (flags.headed
|| flags.executable_path.is_some()
|| flags.state.is_some()
|| flags.proxy.is_some()
|| flags.args.is_some()
|| flags.user_agent.is_some()
|| flags.ignore_https_errors
|| flags.allow_file_access
|| flags.debug
|| flags.color_scheme.is_some()
|| flags.download_path.is_some())
&& flags.cdp.is_none()
&& flags.provider.is_none()
&& !attached_to_existing_browser
{
let mut launch_cmd = json!({
"id": gen_id(),
"action": "launch",
"headless": !flags.headed
});
let cmd_obj = launch_cmd
.as_object_mut()
.expect("json! macro guarantees object type");
// Add executable path if specified
if let Some(ref exec_path) = flags.executable_path {
cmd_obj.insert("executablePath".to_string(), json!(exec_path));
}
// Add state path if specified
if let Some(ref state_path) = flags.state {
cmd_obj.insert("storageState".to_string(), json!(state_path));
}
if let Some(ref proxy_str) = flags.proxy {
let mut proxy_obj = parse_proxy(proxy_str);
// Add bypass if specified
if let Some(ref bypass) = flags.proxy_bypass {
if let Some(obj) = proxy_obj.as_object_mut() {
obj.insert("bypass".to_string(), json!(bypass));
}
}
cmd_obj.insert("proxy".to_string(), proxy_obj);
}
if let Some(ref ua) = flags.user_agent {
cmd_obj.insert("userAgent".to_string(), json!(ua));
}
if let Some(ref a) = flags.args {
// Parse args (comma or newline separated)
let args_vec: Vec<String> = a
.split(&[',', '\n'][..])
.map(|s| s.trim().to_string())
.filter(|s| !s.is_empty())
.collect();
cmd_obj.insert("args".to_string(), json!(args_vec));
}
if flags.ignore_https_errors {
launch_cmd["ignoreHTTPSErrors"] = json!(true);
}
if flags.allow_file_access {
launch_cmd["allowFileAccess"] = json!(true);
}
if let Some(ref cs) = flags.color_scheme {
launch_cmd["colorScheme"] = json!(cs);
}
if let Some(ref dp) = flags.download_path {
launch_cmd["downloadPath"] = json!(dp);
}
if let Some(ref tg) = flags.tab_group {
launch_cmd["tabGroup"] = json!(tg);
}
if let Some(ref plugin_id) = flags.tab_group_plugin_id {
launch_cmd["tabGroupPluginId"] = json!(plugin_id);
}
match send_command(launch_cmd, &flags.session) {
Ok(resp) => {
if !resp.success {
// Launch command failed (e.g., invalid state file)
let error_msg = resp
.error
.unwrap_or_else(|| "Browser launch failed".to_string());
if flags.json {
println!(r#"{{"success":false,"error":"{}"}}"#, error_msg);
} else {
eprintln!("{} {}", color::error_indicator(), error_msg);
}
exit(1);
}
}
Err(e) => {
if flags.json {
println!(r#"{{"success":false,"error":"{}"}}"#, e);
} else {
eprintln!(
"{} Could not configure browser: {}",
color::error_indicator(),
e
);
}
exit(1);
}
}
}
match send_command(cmd.clone(), &flags.session) {
Ok(resp) => {
let success = resp.success;
print_response(&resp, flags.json);
// Extract action for context-specific output handling
let action = cmd.get("action").and_then(|v| v.as_str());
print_response(&resp, flags.json, action);
if !success {
exit(1);
}
@@ -168,9 +796,68 @@ fn main() {
if flags.json {
println!(r#"{{"success":false,"error":"{}"}}"#, e);
} else {
eprintln!("\x1b[31m✗\x1b[0m {}", e);
eprintln!("{} {}", color::error_indicator(), e);
}
exit(1);
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_parse_proxy_simple() {
let result = parse_proxy("http://proxy.com:8080");
assert_eq!(result["server"], "http://proxy.com:8080");
assert!(result.get("username").is_none());
assert!(result.get("password").is_none());
}
#[test]
fn test_parse_proxy_with_auth() {
let result = parse_proxy("http://user:pass@proxy.com:8080");
assert_eq!(result["server"], "http://proxy.com:8080");
assert_eq!(result["username"], "user");
assert_eq!(result["password"], "pass");
}
#[test]
fn test_parse_proxy_username_only() {
let result = parse_proxy("http://user@proxy.com:8080");
assert_eq!(result["server"], "http://proxy.com:8080");
assert_eq!(result["username"], "user");
assert_eq!(result["password"], "");
}
#[test]
fn test_parse_proxy_no_protocol() {
let result = parse_proxy("proxy.com:8080");
assert_eq!(result["server"], "proxy.com:8080");
assert!(result.get("username").is_none());
}
#[test]
fn test_parse_proxy_socks5() {
let result = parse_proxy("socks5://proxy.com:1080");
assert_eq!(result["server"], "socks5://proxy.com:1080");
assert!(result.get("username").is_none());
}
#[test]
fn test_parse_proxy_socks5_with_auth() {
let result = parse_proxy("socks5://admin:secret@proxy.com:1080");
assert_eq!(result["server"], "socks5://proxy.com:1080");
assert_eq!(result["username"], "admin");
assert_eq!(result["password"], "secret");
}
#[test]
fn test_parse_proxy_complex_password() {
let result = parse_proxy("http://user:p@ss:w0rd@proxy.com:8080");
assert_eq!(result["server"], "http://proxy.com:8080");
assert_eq!(result["username"], "user");
assert_eq!(result["password"], "p@ss:w0rd");
}
}
+1
View File
@@ -0,0 +1 @@
include!("main.rs");
File diff suppressed because it is too large Load Diff
+315
View File
@@ -0,0 +1,315 @@
use aes_gcm::{aead::Aead, aead::KeyInit, Aes256Gcm};
use serde::{Deserialize, Serialize};
use serde_json::{json, Value};
use sha2::{Digest, Sha256};
use std::fs;
use std::path::PathBuf;
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct AuthProfile {
pub name: String,
pub url: String,
pub username: String,
pub password: String,
#[serde(skip_serializing_if = "Option::is_none")]
pub username_selector: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub password_selector: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub submit_selector: Option<String>,
}
// Keep legacy Credential alias for backward compatibility
pub type Credential = AuthProfile;
fn validate_profile_name(name: &str) -> Result<(), String> {
if name.is_empty()
|| !name
.chars()
.all(|c| c.is_ascii_alphanumeric() || c == '_' || c == '-')
{
return Err(format!(
"Invalid profile name '{}'. Must match /^[a-zA-Z0-9_-]+$/",
name
));
}
Ok(())
}
fn get_auth_dir() -> PathBuf {
if let Some(home) = dirs::home_dir() {
home.join(".agent-browser").join("auth")
} else {
std::env::temp_dir().join("agent-browser").join("auth")
}
}
fn get_profile_path(name: &str) -> PathBuf {
get_auth_dir().join(format!("{}.json", name))
}
fn derive_encryption_key() -> Vec<u8> {
let hostname = std::env::var("HOSTNAME")
.or_else(|_| std::env::var("COMPUTERNAME"))
.unwrap_or_else(|_| {
#[cfg(unix)]
{
let mut buf = [0u8; 256];
let len = unsafe { libc::gethostname(buf.as_mut_ptr() as *mut _, buf.len()) };
if len == 0 {
let end = buf.iter().position(|&b| b == 0).unwrap_or(buf.len());
String::from_utf8_lossy(&buf[..end]).to_string()
} else {
"unknown-host".to_string()
}
}
#[cfg(not(unix))]
{
"unknown-host".to_string()
}
});
let username = std::env::var("USER")
.or_else(|_| std::env::var("USERNAME"))
.unwrap_or_else(|_| "unknown-user".to_string());
let mut hasher = Sha256::new();
hasher.update(format!("agent-browser:{}:{}", hostname, username).as_bytes());
hasher.finalize().to_vec()
}
fn encrypt_profile(profile: &AuthProfile) -> Result<Vec<u8>, String> {
let key = derive_encryption_key();
let cipher =
Aes256Gcm::new_from_slice(&key).map_err(|e| format!("Encryption key error: {}", e))?;
let plaintext = serde_json::to_string(profile)
.map_err(|e| format!("Failed to serialize profile: {}", e))?;
let mut nonce = [0u8; 12];
getrandom::getrandom(&mut nonce).map_err(|e| format!("Failed to generate nonce: {}", e))?;
let ciphertext = cipher
.encrypt(aes_gcm::Nonce::from_slice(&nonce), plaintext.as_bytes())
.map_err(|e| format!("Encryption failed: {}", e))?;
let mut result = Vec::with_capacity(12 + ciphertext.len());
result.extend_from_slice(&nonce);
result.extend_from_slice(&ciphertext);
Ok(result)
}
fn decrypt_profile(data: &[u8]) -> Result<AuthProfile, String> {
if data.len() < 13 {
return Err("Encrypted data too short".to_string());
}
let (nonce_bytes, ciphertext) = data.split_at(12);
let key = derive_encryption_key();
let cipher =
Aes256Gcm::new_from_slice(&key).map_err(|e| format!("Decryption key error: {}", e))?;
let plaintext = cipher
.decrypt(aes_gcm::Nonce::from_slice(nonce_bytes), ciphertext)
.map_err(|e| format!("Decryption failed: {}", e))?;
let json_str = String::from_utf8(plaintext)
.map_err(|e| format!("Decrypted data is not valid UTF-8: {}", e))?;
serde_json::from_str(&json_str).map_err(|e| format!("Invalid profile data: {}", e))
}
fn save_profile(profile: &AuthProfile) -> Result<(), String> {
let dir = get_auth_dir();
let _ = fs::create_dir_all(&dir);
let encrypted = encrypt_profile(profile)?;
let path = get_profile_path(&profile.name);
fs::write(&path, &encrypted).map_err(|e| format!("Failed to write profile: {}", e))
}
fn load_profile(name: &str) -> Result<AuthProfile, String> {
let path = get_profile_path(name);
if !path.exists() {
return Err(format!("Auth profile '{}' not found", name));
}
let data = fs::read(&path).map_err(|e| format!("Failed to read profile: {}", e))?;
decrypt_profile(&data)
}
pub fn credentials_set(
name: &str,
username: &str,
password: &str,
url: Option<&str>,
) -> Result<Value, String> {
validate_profile_name(name)?;
let profile = AuthProfile {
name: name.to_string(),
url: url.unwrap_or("").to_string(),
username: username.to_string(),
password: password.to_string(),
username_selector: None,
password_selector: None,
submit_selector: None,
};
save_profile(&profile)?;
Ok(json!({ "saved": name }))
}
pub fn auth_save(
name: &str,
url: &str,
username: &str,
password: &str,
username_selector: Option<&str>,
password_selector: Option<&str>,
submit_selector: Option<&str>,
) -> Result<Value, String> {
validate_profile_name(name)?;
let profile = AuthProfile {
name: name.to_string(),
url: url.to_string(),
username: username.to_string(),
password: password.to_string(),
username_selector: username_selector.map(String::from),
password_selector: password_selector.map(String::from),
submit_selector: submit_selector.map(String::from),
};
save_profile(&profile)?;
Ok(json!({ "saved": name }))
}
pub fn credentials_get(name: &str) -> Result<Value, String> {
let profile = load_profile(name)?;
Ok(json!({
"name": profile.name,
"username": profile.username,
"url": profile.url,
"hasPassword": true,
}))
}
pub fn credentials_get_full(name: &str) -> Result<AuthProfile, String> {
load_profile(name)
}
pub fn credentials_delete(name: &str) -> Result<Value, String> {
validate_profile_name(name)?;
let path = get_profile_path(name);
if !path.exists() {
return Err(format!("Auth profile '{}' not found", name));
}
fs::remove_file(&path).map_err(|e| format!("Failed to delete profile: {}", e))?;
Ok(json!({ "deleted": name }))
}
pub fn credentials_list() -> Result<Value, String> {
let dir = get_auth_dir();
if !dir.exists() {
return Ok(json!({ "profiles": [] }));
}
let mut profiles = Vec::new();
if let Ok(entries) = fs::read_dir(&dir) {
for entry in entries.flatten() {
let path = entry.path();
if path.extension().and_then(|e| e.to_str()) != Some("json") {
continue;
}
let name = path
.file_stem()
.unwrap_or_default()
.to_string_lossy()
.to_string();
match load_profile(&name) {
Ok(profile) => {
profiles.push(json!({
"name": profile.name,
"username": profile.username,
"url": profile.url,
}));
}
Err(_) => {
profiles.push(json!({
"name": name,
"error": "Failed to decrypt",
}));
}
}
}
}
Ok(json!({ "profiles": profiles }))
}
pub fn auth_show(name: &str) -> Result<Value, String> {
validate_profile_name(name)?;
let profile = load_profile(name)?;
Ok(json!({
"profile": {
"name": profile.name,
"url": profile.url,
"username": profile.username,
"usernameSelector": profile.username_selector,
"passwordSelector": profile.password_selector,
"submitSelector": profile.submit_selector,
}
}))
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_validate_profile_name() {
assert!(validate_profile_name("github").is_ok());
assert!(validate_profile_name("my-app").is_ok());
assert!(validate_profile_name("test_123").is_ok());
assert!(validate_profile_name("").is_err());
assert!(validate_profile_name("has space").is_err());
assert!(validate_profile_name("../evil").is_err());
assert!(validate_profile_name("foo/bar").is_err());
}
#[test]
fn test_auth_profile_serialization() {
let profile = AuthProfile {
name: "test".to_string(),
url: "https://example.com".to_string(),
username: "user".to_string(),
password: "pass".to_string(),
username_selector: None,
password_selector: None,
submit_selector: Some("button[type=submit]".to_string()),
};
let json = serde_json::to_string(&profile).unwrap();
let parsed: AuthProfile = serde_json::from_str(&json).unwrap();
assert_eq!(parsed.name, "test");
assert_eq!(
parsed.submit_selector,
Some("button[type=submit]".to_string())
);
assert!(parsed.username_selector.is_none());
}
#[test]
fn test_encrypt_decrypt_roundtrip() {
let profile = AuthProfile {
name: "roundtrip".to_string(),
url: "https://example.com".to_string(),
username: "user".to_string(),
password: "s3cret!".to_string(),
username_selector: None,
password_selector: None,
submit_selector: None,
};
let encrypted = encrypt_profile(&profile).unwrap();
let decrypted = decrypt_profile(&encrypted).unwrap();
assert_eq!(decrypted.name, "roundtrip");
assert_eq!(decrypted.password, "s3cret!");
}
#[test]
fn test_derive_encryption_key_is_stable() {
let k1 = derive_encryption_key();
let k2 = derive_encryption_key();
assert_eq!(k1, k2);
assert_eq!(k1.len(), 32);
}
}
File diff suppressed because it is too large Load Diff
+783
View File
@@ -0,0 +1,783 @@
use std::io::{BufRead, BufReader};
use std::path::{Path, PathBuf};
use std::process::{Child, Command, Stdio};
use std::time::Duration;
use super::types::BrowserVersionInfo;
pub struct ChromeProcess {
child: Child,
pub ws_url: String,
temp_user_data_dir: Option<PathBuf>,
}
impl ChromeProcess {
pub fn kill(&mut self) {
let _ = self.child.kill();
let _ = self.child.wait();
}
}
impl Drop for ChromeProcess {
fn drop(&mut self) {
self.kill();
if let Some(ref dir) = self.temp_user_data_dir {
for attempt in 0..3 {
match std::fs::remove_dir_all(dir) {
Ok(()) => break,
Err(_) if attempt < 2 => {
std::thread::sleep(Duration::from_millis(100));
}
Err(e) => {
eprintln!(
"Warning: failed to clean up temp profile {}: {}",
dir.display(),
e
);
}
}
}
}
}
}
pub struct LaunchOptions {
pub headless: bool,
pub executable_path: Option<String>,
pub proxy: Option<String>,
pub proxy_bypass: Option<String>,
pub profile: Option<String>,
pub args: Vec<String>,
pub allow_file_access: bool,
pub extensions: Option<Vec<String>>,
pub storage_state: Option<String>,
pub user_agent: Option<String>,
pub ignore_https_errors: bool,
pub color_scheme: Option<String>,
pub download_path: Option<String>,
}
impl Default for LaunchOptions {
fn default() -> Self {
Self {
headless: true,
executable_path: None,
proxy: None,
proxy_bypass: None,
profile: None,
args: Vec::new(),
allow_file_access: false,
extensions: None,
storage_state: None,
user_agent: None,
ignore_https_errors: false,
color_scheme: None,
download_path: None,
}
}
}
struct ChromeArgs {
args: Vec<String>,
temp_user_data_dir: Option<PathBuf>,
}
fn build_chrome_args(options: &LaunchOptions) -> Result<ChromeArgs, String> {
let mut args = vec![
"--remote-debugging-port=0".to_string(),
"--no-first-run".to_string(),
"--no-default-browser-check".to_string(),
"--disable-background-networking".to_string(),
"--disable-backgrounding-occluded-windows".to_string(),
"--disable-component-update".to_string(),
"--disable-default-apps".to_string(),
"--disable-hang-monitor".to_string(),
"--disable-popup-blocking".to_string(),
"--disable-prompt-on-repost".to_string(),
"--disable-sync".to_string(),
"--enable-features=NetworkService,NetworkServiceInProcess".to_string(),
"--metrics-recording-only".to_string(),
"--password-store=basic".to_string(),
"--use-mock-keychain".to_string(),
];
if options.headless {
args.push("--headless=new".to_string());
}
if let Some(ref proxy) = options.proxy {
args.push(format!("--proxy-server={}", proxy));
}
if let Some(ref bypass) = options.proxy_bypass {
args.push(format!("--proxy-bypass-list={}", bypass));
}
let temp_user_data_dir = if let Some(ref profile) = options.profile {
let expanded = expand_tilde(profile);
args.push(format!("--user-data-dir={}", expanded));
None
} else {
let dir = std::env::temp_dir()
.join(format!("agent-browser-chrome-{}", uuid::Uuid::new_v4()));
std::fs::create_dir_all(&dir)
.map_err(|e| format!("Failed to create temp profile dir: {}", e))?;
args.push(format!("--user-data-dir={}", dir.display()));
Some(dir)
};
if options.allow_file_access {
args.push("--allow-file-access-from-files".to_string());
args.push("--allow-file-access".to_string());
}
if let Some(ref exts) = options.extensions {
if !exts.is_empty() {
let ext_list = exts.join(",");
args.push(format!("--load-extension={}", ext_list));
args.push(format!("--disable-extensions-except={}", ext_list));
}
}
let has_window_size = options
.args
.iter()
.any(|a| a.starts_with("--start-maximized") || a.starts_with("--window-size="));
if !has_window_size && options.headless {
args.push("--window-size=1280,720".to_string());
}
args.extend(options.args.iter().cloned());
if should_disable_sandbox(&args) {
args.push("--no-sandbox".to_string());
}
Ok(ChromeArgs {
args,
temp_user_data_dir,
})
}
pub fn launch_chrome(options: &LaunchOptions) -> Result<ChromeProcess, String> {
let chrome_path = match &options.executable_path {
Some(p) => PathBuf::from(p),
None => {
find_chrome().ok_or("Chrome not found. Install Chrome or use --executable-path.")?
}
};
let ChromeArgs {
args,
temp_user_data_dir,
} = build_chrome_args(options)?;
let cleanup_temp_dir = |dir: &Option<PathBuf>| {
if let Some(ref d) = dir {
let _ = std::fs::remove_dir_all(d);
}
};
let mut child = Command::new(&chrome_path)
.args(&args)
.stdin(Stdio::null())
.stdout(Stdio::null())
.stderr(Stdio::piped())
.spawn()
.map_err(|e| {
cleanup_temp_dir(&temp_user_data_dir);
format!("Failed to launch Chrome at {:?}: {}", chrome_path, e)
})?;
let stderr = child
.stderr
.take()
.ok_or_else(|| {
let _ = child.kill();
cleanup_temp_dir(&temp_user_data_dir);
"Failed to capture Chrome stderr".to_string()
})?;
let reader = BufReader::new(stderr);
let ws_url = match wait_for_ws_url(reader) {
Ok(url) => url,
Err(e) => {
let _ = child.kill();
cleanup_temp_dir(&temp_user_data_dir);
return Err(e);
}
};
Ok(ChromeProcess {
child,
ws_url,
temp_user_data_dir,
})
}
fn wait_for_ws_url(reader: BufReader<std::process::ChildStderr>) -> Result<String, String> {
let deadline = std::time::Instant::now() + Duration::from_secs(30);
let prefix = "DevTools listening on ";
let mut stderr_lines: Vec<String> = Vec::new();
for line in reader.lines() {
if std::time::Instant::now() > deadline {
return Err(chrome_launch_error(
"Timeout waiting for Chrome DevTools URL",
&stderr_lines,
));
}
let line = line.map_err(|e| format!("Failed to read Chrome stderr: {}", e))?;
if let Some(url) = line.strip_prefix(prefix) {
return Ok(url.trim().to_string());
}
stderr_lines.push(line);
}
Err(chrome_launch_error(
"Chrome exited before providing DevTools URL",
&stderr_lines,
))
}
fn chrome_launch_error(message: &str, stderr_lines: &[String]) -> String {
let relevant: Vec<&String> = stderr_lines
.iter()
.filter(|l| {
let lower = l.to_lowercase();
lower.contains("error")
|| lower.contains("fatal")
|| lower.contains("sandbox")
|| lower.contains("namespace")
|| lower.contains("permission")
|| lower.contains("cannot")
|| lower.contains("failed")
|| lower.contains("abort")
})
.collect();
if relevant.is_empty() {
if stderr_lines.is_empty() {
return format!("{} (no stderr output from Chrome)", message);
}
let last_lines: Vec<&String> = stderr_lines.iter().rev().take(5).collect();
return format!(
"{}\nChrome stderr (last {} lines):\n {}",
message,
last_lines.len(),
last_lines
.into_iter()
.rev()
.map(|s| s.as_str())
.collect::<Vec<_>>()
.join("\n ")
);
}
let hint = if relevant.iter().any(|l| {
let lower = l.to_lowercase();
lower.contains("sandbox") || lower.contains("namespace")
}) {
"\nHint: try --args \"--no-sandbox\" (required in containers, VMs, and some Linux setups)"
} else {
""
};
format!(
"{}\nChrome stderr:\n {}{}",
message,
relevant
.iter()
.map(|s| s.as_str())
.collect::<Vec<_>>()
.join("\n "),
hint
)
}
pub fn find_chrome() -> Option<PathBuf> {
#[cfg(target_os = "macos")]
{
let candidates = [
"/Applications/Google Chrome.app/Contents/MacOS/Google Chrome",
"/Applications/Google Chrome Canary.app/Contents/MacOS/Google Chrome Canary",
"/Applications/Chromium.app/Contents/MacOS/Chromium",
];
for c in &candidates {
let p = PathBuf::from(c);
if p.exists() {
return Some(p);
}
}
if let Some(p) = find_playwright_chromium() {
return Some(p);
}
}
#[cfg(target_os = "linux")]
{
let candidates = [
"google-chrome",
"google-chrome-stable",
"chromium-browser",
"chromium",
];
for name in &candidates {
if let Ok(output) = Command::new("which").arg(name).output() {
if output.status.success() {
let path = String::from_utf8_lossy(&output.stdout).trim().to_string();
if !path.is_empty() {
return Some(PathBuf::from(path));
}
}
}
}
if let Some(p) = find_playwright_chromium() {
return Some(p);
}
}
#[cfg(target_os = "windows")]
{
let candidates = [
r"C:\Program Files\Google\Chrome\Application\chrome.exe",
r"C:\Program Files (x86)\Google\Chrome\Application\chrome.exe",
];
if let Ok(local) = std::env::var("LOCALAPPDATA") {
let p = PathBuf::from(&local).join(r"Google\Chrome\Application\chrome.exe");
if p.exists() {
return Some(p);
}
}
for c in &candidates {
let p = PathBuf::from(c);
if p.exists() {
return Some(p);
}
}
}
None
}
pub async fn discover_cdp_url(port: u16) -> Result<String, String> {
let url = format!("http://127.0.0.1:{}/json/version", port);
let body = tokio::time::timeout(Duration::from_secs(2), async {
reqwest_get_string(&url).await
})
.await
.map_err(|_| format!("Timeout connecting to CDP on port {}", port))?
.map_err(|e| format!("Failed to connect to CDP on port {}: {}", port, e))?;
let info: BrowserVersionInfo = serde_json::from_str(&body)
.map_err(|e| format!("Invalid /json/version response: {}", e))?;
info.web_socket_debugger_url
.ok_or_else(|| format!("No webSocketDebuggerUrl in /json/version on port {}", port))
}
async fn reqwest_get_string(url: &str) -> Result<String, String> {
let resp = reqwest::get(url).await.map_err(|e| e.to_string())?;
resp.text().await.map_err(|e| e.to_string())
}
pub fn read_devtools_active_port(user_data_dir: &Path) -> Option<(u16, String)> {
let path = user_data_dir.join("DevToolsActivePort");
let content = std::fs::read_to_string(&path).ok()?;
let mut lines = content.lines();
let port: u16 = lines.next()?.trim().parse().ok()?;
let ws_path = lines
.next()
.unwrap_or("/devtools/browser")
.trim()
.to_string();
Some((port, ws_path))
}
pub async fn auto_connect_cdp() -> Result<String, String> {
let user_data_dirs = get_chrome_user_data_dirs();
for dir in &user_data_dirs {
if let Some((port, ws_path)) = read_devtools_active_port(dir) {
// Try HTTP endpoint first (pre-M144)
if let Ok(ws_url) = discover_cdp_url(port).await {
return Ok(ws_url);
}
// M144+: direct WebSocket
let ws_url = format!("ws://127.0.0.1:{}{}", port, ws_path);
return Ok(ws_url);
}
}
// Fallback: probe common ports
for port in [9222u16, 9229] {
if let Ok(ws_url) = discover_cdp_url(port).await {
return Ok(ws_url);
}
}
Err("No running Chrome instance found. Launch Chrome with --remote-debugging-port or use --cdp.".to_string())
}
fn get_chrome_user_data_dirs() -> Vec<PathBuf> {
let mut dirs = Vec::new();
#[cfg(target_os = "macos")]
{
if let Some(home) = dirs::home_dir() {
let base = home.join("Library/Application Support");
for name in ["Google/Chrome", "Google/Chrome Canary", "Chromium"] {
dirs.push(base.join(name));
}
}
}
#[cfg(target_os = "linux")]
{
if let Some(home) = dirs::home_dir() {
let config = home.join(".config");
for name in ["google-chrome", "google-chrome-unstable", "chromium"] {
dirs.push(config.join(name));
}
}
}
#[cfg(target_os = "windows")]
{
if let Ok(local) = std::env::var("LOCALAPPDATA") {
let base = PathBuf::from(local);
for name in [
r"Google\Chrome\User Data",
r"Google\Chrome SxS\User Data",
r"Chromium\User Data",
] {
dirs.push(base.join(name));
}
}
}
dirs
}
/// Returns true if Chrome's sandbox should be disabled because the environment
/// doesn't support it (containers, VMs, running as root).
fn should_disable_sandbox(existing_args: &[String]) -> bool {
if existing_args.iter().any(|a| a == "--no-sandbox") {
return false; // already set by user
}
#[cfg(unix)]
{
// Root user -- standard container default, Chrome sandbox requires non-root
if unsafe { libc::geteuid() } == 0 {
return true;
}
// Docker container
if Path::new("/.dockerenv").exists() {
return true;
}
// Podman container
if Path::new("/run/.containerenv").exists() {
return true;
}
// Generic container detection: cgroup contains docker/kubepods/lxc
if let Ok(cgroup) = std::fs::read_to_string("/proc/1/cgroup") {
if cgroup.contains("docker")
|| cgroup.contains("kubepods")
|| cgroup.contains("lxc")
{
return true;
}
}
}
false
}
/// Search Playwright's browser cache for a Chromium binary.
/// This is where `agent-browser install` (via `npx playwright install chromium`) puts it.
fn find_playwright_chromium() -> Option<PathBuf> {
let mut search_dirs = Vec::new();
if let Ok(custom) = std::env::var("PLAYWRIGHT_BROWSERS_PATH") {
search_dirs.push(PathBuf::from(custom));
}
if let Some(home) = dirs::home_dir() {
search_dirs.push(home.join(".cache/ms-playwright"));
}
for dir in &search_dirs {
if !dir.is_dir() {
continue;
}
if let Ok(entries) = std::fs::read_dir(dir) {
let mut matches: Vec<PathBuf> = entries
.filter_map(|e| e.ok())
.filter(|e| {
e.file_name()
.to_str()
.map(|n| n.starts_with("chromium-"))
.unwrap_or(false)
})
.filter_map(|e| {
let candidate = build_playwright_binary_path(&e.path());
if candidate.exists() {
Some(candidate)
} else {
None
}
})
.collect();
// Sort descending so the newest version wins
matches.sort();
matches.reverse();
if let Some(p) = matches.into_iter().next() {
return Some(p);
}
}
}
None
}
#[cfg(target_os = "linux")]
fn build_playwright_binary_path(chromium_dir: &Path) -> PathBuf {
chromium_dir.join("chrome-linux64/chrome")
}
#[cfg(target_os = "macos")]
fn build_playwright_binary_path(chromium_dir: &Path) -> PathBuf {
chromium_dir.join("chrome-mac/Chromium.app/Contents/MacOS/Chromium")
}
#[cfg(target_os = "windows")]
fn build_playwright_binary_path(chromium_dir: &Path) -> PathBuf {
chromium_dir.join("chrome-win/chrome.exe")
}
fn expand_tilde(path: &str) -> String {
if let Some(rest) = path.strip_prefix('~') {
if let Some(home) = dirs::home_dir() {
return home
.join(rest.strip_prefix('/').unwrap_or(rest))
.to_string_lossy()
.to_string();
}
}
path.to_string()
}
#[cfg(test)]
mod tests {
use super::*;
use crate::test_utils::EnvGuard;
#[test]
fn test_find_chrome_returns_some_on_host() {
// This test only makes sense on systems with Chrome installed
if cfg!(target_os = "macos") || cfg!(target_os = "linux") {
let result = find_chrome();
// Don't assert Some -- CI may not have Chrome
if let Some(path) = result {
assert!(path.exists());
}
}
}
#[test]
fn test_expand_tilde() {
let expanded = expand_tilde("~/test/path");
assert!(!expanded.starts_with('~'));
assert!(expanded.ends_with("test/path"));
}
#[test]
fn test_expand_tilde_no_tilde() {
assert_eq!(expand_tilde("/absolute/path"), "/absolute/path");
}
#[test]
fn test_read_devtools_active_port_missing() {
let result = read_devtools_active_port(Path::new("/nonexistent"));
assert!(result.is_none());
}
#[test]
fn test_should_disable_sandbox_skips_if_already_set() {
let args = vec!["--headless=new".to_string(), "--no-sandbox".to_string()];
assert!(!should_disable_sandbox(&args));
}
#[test]
fn test_chrome_launch_error_no_stderr() {
let msg = chrome_launch_error("Chrome exited", &[]);
assert!(msg.contains("no stderr output"));
}
#[test]
fn test_chrome_launch_error_with_sandbox_hint() {
let lines = vec![
"some log line".to_string(),
"Failed to move to new namespace: sandbox error".to_string(),
];
let msg = chrome_launch_error("Chrome exited", &lines);
assert!(msg.contains("sandbox error"));
assert!(msg.contains("Hint:"));
assert!(msg.contains("--no-sandbox"));
}
#[test]
fn test_chrome_launch_error_generic() {
let lines = vec![
"info line".to_string(),
"another info line".to_string(),
];
let msg = chrome_launch_error("Chrome exited", &lines);
assert!(msg.contains("last 2 lines"));
}
#[test]
fn test_find_playwright_chromium_nonexistent() {
let _guard = EnvGuard::new(&["PLAYWRIGHT_BROWSERS_PATH"]);
_guard.set("PLAYWRIGHT_BROWSERS_PATH", "/nonexistent/path");
let result = find_playwright_chromium();
assert!(result.is_none());
}
#[test]
fn test_build_args_headless_includes_headless_flag() {
let opts = LaunchOptions {
headless: true,
..Default::default()
};
let result = build_chrome_args(&opts).unwrap();
assert!(result.args.iter().any(|a| a == "--headless=new"));
assert!(result
.args
.iter()
.any(|a| a == "--window-size=1280,720"));
// Temp dir created when no profile
assert!(result.temp_user_data_dir.is_some());
let dir = result.temp_user_data_dir.unwrap();
assert!(dir.exists());
let _ = std::fs::remove_dir_all(&dir);
}
#[test]
fn test_build_args_headed_no_headless_flag() {
let opts = LaunchOptions {
headless: false,
..Default::default()
};
let result = build_chrome_args(&opts).unwrap();
assert!(!result.args.iter().any(|a| a.contains("--headless")));
assert!(!result.args.iter().any(|a| a.starts_with("--window-size=")));
// Temp dir created when no profile
assert!(result.temp_user_data_dir.is_some());
let dir = result.temp_user_data_dir.unwrap();
assert!(dir.exists());
let _ = std::fs::remove_dir_all(&dir);
}
#[test]
fn test_build_args_temp_user_data_dir_created() {
let opts = LaunchOptions::default();
let result = build_chrome_args(&opts).unwrap();
let dir = result.temp_user_data_dir.as_ref().unwrap();
assert!(dir.exists());
assert!(result
.args
.iter()
.any(|a| a.starts_with("--user-data-dir=")));
let _ = std::fs::remove_dir_all(dir);
}
#[test]
fn test_build_args_profile_no_temp_dir() {
let opts = LaunchOptions {
profile: Some("/tmp/my-profile".to_string()),
..Default::default()
};
let result = build_chrome_args(&opts).unwrap();
assert!(result.temp_user_data_dir.is_none());
assert!(result
.args
.iter()
.any(|a| a == "--user-data-dir=/tmp/my-profile"));
}
#[test]
fn test_build_args_custom_window_size_not_overridden() {
let opts = LaunchOptions {
headless: true,
args: vec!["--window-size=1920,1080".to_string()],
..Default::default()
};
let result = build_chrome_args(&opts).unwrap();
assert!(!result
.args
.iter()
.any(|a| a == "--window-size=1280,720"));
assert!(result
.args
.iter()
.any(|a| a == "--window-size=1920,1080"));
if let Some(ref dir) = result.temp_user_data_dir {
let _ = std::fs::remove_dir_all(dir);
}
}
#[test]
fn test_build_args_start_maximized_suppresses_default_window_size() {
let opts = LaunchOptions {
headless: true,
args: vec!["--start-maximized".to_string()],
..Default::default()
};
let result = build_chrome_args(&opts).unwrap();
assert!(!result.args.iter().any(|a| a == "--window-size=1280,720"));
assert!(result.args.iter().any(|a| a == "--start-maximized"));
if let Some(ref dir) = result.temp_user_data_dir {
let _ = std::fs::remove_dir_all(dir);
}
}
#[test]
fn test_chrome_process_drop_cleans_temp_dir() {
let dir = std::env::temp_dir().join(format!(
"agent-browser-chrome-drop-test-{}",
uuid::Uuid::new_v4()
));
let _ = std::fs::create_dir_all(&dir);
assert!(dir.exists());
{
// Simulate a ChromeProcess with a temp dir but a dummy child.
// We can't actually spawn Chrome here, but we can verify the Drop
// logic by creating a small helper process.
let child = Command::new("echo")
.arg("test")
.stdin(Stdio::null())
.stdout(Stdio::null())
.stderr(Stdio::null())
.spawn()
.unwrap();
let _process = ChromeProcess {
child,
ws_url: String::new(),
temp_user_data_dir: Some(dir.clone()),
};
// _process dropped here
}
assert!(!dir.exists(), "Temp dir should be cleaned up on drop");
}
}
+163
View File
@@ -0,0 +1,163 @@
use std::collections::HashMap;
use std::sync::atomic::{AtomicU64, Ordering};
use std::sync::Arc;
use futures_util::{SinkExt, StreamExt};
use serde_json::Value;
use tokio::sync::{broadcast, oneshot, Mutex};
use tokio_tungstenite::connect_async;
use tokio_tungstenite::tungstenite::Message;
use super::types::{CdpCommand, CdpEvent, CdpMessage};
type PendingMap = Arc<Mutex<HashMap<u64, oneshot::Sender<CdpMessage>>>>;
pub struct CdpClient {
ws_tx: Arc<
Mutex<
futures_util::stream::SplitSink<
tokio_tungstenite::WebSocketStream<
tokio_tungstenite::MaybeTlsStream<tokio::net::TcpStream>,
>,
Message,
>,
>,
>,
next_id: AtomicU64,
pending: PendingMap,
event_tx: broadcast::Sender<CdpEvent>,
_reader_handle: tokio::task::JoinHandle<()>,
}
impl CdpClient {
pub async fn connect(url: &str) -> Result<Self, String> {
let (ws_stream, _) = connect_async(url)
.await
.map_err(|e| format!("CDP WebSocket connect failed: {}", e))?;
let (ws_tx, mut ws_rx) = ws_stream.split();
let ws_tx = Arc::new(Mutex::new(ws_tx));
let pending: PendingMap = Arc::new(Mutex::new(HashMap::new()));
let (event_tx, _) = broadcast::channel(256);
let pending_clone = pending.clone();
let event_tx_clone = event_tx.clone();
let reader_handle = tokio::spawn(async move {
while let Some(msg) = ws_rx.next().await {
let msg = match msg {
Ok(Message::Text(text)) => text,
Ok(Message::Close(_)) => break,
Ok(_) => continue,
Err(_) => break,
};
let parsed: CdpMessage = match serde_json::from_str(&msg) {
Ok(m) => m,
Err(_) => continue,
};
if let Some(id) = parsed.id {
// Response to a command
let mut pending = pending_clone.lock().await;
if let Some(tx) = pending.remove(&id) {
let _ = tx.send(parsed);
}
} else if let Some(ref method) = parsed.method {
// Event
let event = CdpEvent {
method: method.clone(),
params: parsed.params.clone().unwrap_or(Value::Null),
session_id: parsed.session_id.clone(),
};
let _ = event_tx_clone.send(event);
}
}
});
Ok(Self {
ws_tx,
next_id: AtomicU64::new(1),
pending,
event_tx,
_reader_handle: reader_handle,
})
}
pub async fn send_command(
&self,
method: &str,
params: Option<Value>,
session_id: Option<&str>,
) -> Result<Value, String> {
let id = self.next_id.fetch_add(1, Ordering::SeqCst);
let cmd = CdpCommand {
id,
method: method.to_string(),
params,
session_id: session_id.map(|s| s.to_string()),
};
let json = serde_json::to_string(&cmd)
.map_err(|e| format!("Failed to serialize CDP command: {}", e))?;
let (tx, rx) = oneshot::channel();
{
let mut pending = self.pending.lock().await;
pending.insert(id, tx);
}
{
let mut ws_tx = self.ws_tx.lock().await;
ws_tx
.send(Message::Text(json))
.await
.map_err(|e| format!("Failed to send CDP command: {}", e))?;
}
let response = match tokio::time::timeout(std::time::Duration::from_secs(30), rx).await {
Ok(Ok(resp)) => resp,
Ok(Err(_)) => return Err("CDP response channel closed".to_string()),
Err(_) => {
self.pending.lock().await.remove(&id);
return Err(format!("CDP command timed out: {}", method));
}
};
if let Some(error) = response.error {
return Err(format!("CDP error ({}): {}", method, error));
}
Ok(response.result.unwrap_or(Value::Null))
}
pub fn subscribe(&self) -> broadcast::Receiver<CdpEvent> {
self.event_tx.subscribe()
}
pub async fn send_command_typed<P: serde::Serialize, R: serde::de::DeserializeOwned>(
&self,
method: &str,
params: &P,
session_id: Option<&str>,
) -> Result<R, String> {
let params_value = serde_json::to_value(params)
.map_err(|e| format!("Failed to serialize params: {}", e))?;
let result = self
.send_command(method, Some(params_value), session_id)
.await?;
serde_json::from_value(result)
.map_err(|e| format!("Failed to deserialize CDP response for {}: {}", method, e))
}
pub async fn send_command_no_params(
&self,
method: &str,
session_id: Option<&str>,
) -> Result<Value, String> {
self.send_command(method, None, session_id).await
}
}
+3
View File
@@ -0,0 +1,3 @@
pub mod chrome;
pub mod client;
pub mod types;
+537
View File
@@ -0,0 +1,537 @@
use serde::{Deserialize, Serialize};
use serde_json::Value;
// ---------------------------------------------------------------------------
// CDP message envelope
// ---------------------------------------------------------------------------
#[derive(Debug, Serialize)]
#[serde(rename_all = "camelCase")]
pub struct CdpCommand {
pub id: u64,
pub method: String,
#[serde(skip_serializing_if = "Option::is_none")]
pub params: Option<Value>,
#[serde(skip_serializing_if = "Option::is_none")]
pub session_id: Option<String>,
}
#[derive(Debug, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct CdpMessage {
pub id: Option<u64>,
pub result: Option<Value>,
pub error: Option<CdpError>,
pub method: Option<String>,
pub params: Option<Value>,
pub session_id: Option<String>,
}
#[derive(Debug, Clone, Deserialize)]
pub struct CdpError {
pub code: Option<i64>,
pub message: String,
pub data: Option<String>,
}
impl std::fmt::Display for CdpError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(f, "{}", self.message)
}
}
// ---------------------------------------------------------------------------
// CDP events (broadcast to subscribers)
// ---------------------------------------------------------------------------
#[derive(Debug, Clone)]
pub struct CdpEvent {
pub method: String,
pub params: Value,
pub session_id: Option<String>,
}
// ---------------------------------------------------------------------------
// Target domain
// ---------------------------------------------------------------------------
#[derive(Debug, Clone, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct TargetInfo {
pub target_id: String,
#[serde(rename = "type")]
pub target_type: String,
pub title: String,
pub url: String,
pub attached: Option<bool>,
pub browser_context_id: Option<String>,
}
#[derive(Debug, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct GetTargetsResult {
pub target_infos: Vec<TargetInfo>,
}
#[derive(Debug, Serialize)]
#[serde(rename_all = "camelCase")]
pub struct AttachToTargetParams {
pub target_id: String,
pub flatten: bool,
}
#[derive(Debug, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct AttachToTargetResult {
pub session_id: String,
}
#[derive(Debug, Serialize)]
#[serde(rename_all = "camelCase")]
pub struct SetDiscoverTargetsParams {
pub discover: bool,
}
#[derive(Debug, Serialize)]
#[serde(rename_all = "camelCase")]
pub struct CreateTargetParams {
pub url: String,
}
#[derive(Debug, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct CreateTargetResult {
pub target_id: String,
}
#[derive(Debug, Serialize)]
#[serde(rename_all = "camelCase")]
pub struct CloseTargetParams {
pub target_id: String,
}
// Target events
#[derive(Debug, Clone, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct TargetCreatedEvent {
pub target_info: TargetInfo,
}
#[derive(Debug, Clone, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct TargetDestroyedEvent {
pub target_id: String,
}
#[derive(Debug, Clone, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct TargetInfoChangedEvent {
pub target_info: TargetInfo,
}
// ---------------------------------------------------------------------------
// Page domain
// ---------------------------------------------------------------------------
#[derive(Debug, Serialize)]
#[serde(rename_all = "camelCase")]
pub struct PageNavigateParams {
pub url: String,
#[serde(skip_serializing_if = "Option::is_none")]
pub referrer: Option<String>,
}
#[derive(Debug, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct PageNavigateResult {
pub frame_id: String,
pub loader_id: Option<String>,
pub error_text: Option<String>,
}
#[derive(Debug, Clone, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct FrameNavigatedEvent {
pub frame: FrameInfo,
}
#[derive(Debug, Clone, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct FrameInfo {
pub id: String,
pub url: String,
pub parent_id: Option<String>,
pub name: Option<String>,
}
// Page.javascriptDialogOpening
#[derive(Debug, Clone, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct JavascriptDialogOpeningEvent {
pub url: String,
pub message: String,
#[serde(rename = "type")]
pub dialog_type: String,
pub default_prompt: Option<String>,
}
#[derive(Debug, Serialize)]
#[serde(rename_all = "camelCase")]
pub struct HandleJavaScriptDialogParams {
pub accept: bool,
#[serde(skip_serializing_if = "Option::is_none")]
pub prompt_text: Option<String>,
}
// ---------------------------------------------------------------------------
// Runtime domain
// ---------------------------------------------------------------------------
#[derive(Debug, Serialize)]
#[serde(rename_all = "camelCase")]
pub struct EvaluateParams {
pub expression: String,
#[serde(skip_serializing_if = "Option::is_none")]
pub return_by_value: Option<bool>,
#[serde(skip_serializing_if = "Option::is_none")]
pub await_promise: Option<bool>,
}
#[derive(Debug, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct EvaluateResult {
pub result: RemoteObject,
pub exception_details: Option<ExceptionDetails>,
}
#[derive(Debug, Clone, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct RemoteObject {
#[serde(rename = "type")]
pub object_type: String,
pub subtype: Option<String>,
pub value: Option<Value>,
pub description: Option<String>,
pub object_id: Option<String>,
pub class_name: Option<String>,
pub unserializable_value: Option<String>,
}
#[derive(Debug, Clone, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct ExceptionDetails {
pub text: String,
pub exception: Option<RemoteObject>,
pub line_number: Option<i64>,
pub column_number: Option<i64>,
}
// Runtime.consoleAPICalled
#[derive(Debug, Clone, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct ConsoleApiCalledEvent {
#[serde(rename = "type")]
pub call_type: String,
pub args: Vec<RemoteObject>,
pub timestamp: Option<f64>,
}
// Runtime.exceptionThrown
#[derive(Debug, Clone, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct ExceptionThrownEvent {
pub timestamp: f64,
pub exception_details: ExceptionDetails,
}
// ---------------------------------------------------------------------------
// Accessibility domain
// ---------------------------------------------------------------------------
#[derive(Debug, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct GetFullAXTreeResult {
pub nodes: Vec<AXNode>,
}
#[derive(Debug, Clone, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct AXNode {
pub node_id: String,
pub role: Option<AXValue>,
pub name: Option<AXValue>,
pub value: Option<AXValue>,
pub description: Option<AXValue>,
pub properties: Option<Vec<AXProperty>>,
pub child_ids: Option<Vec<String>>,
pub backend_d_o_m_node_id: Option<i64>,
pub ignored: Option<bool>,
}
#[derive(Debug, Clone, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct AXValue {
#[serde(rename = "type")]
pub value_type: String,
pub value: Option<Value>,
}
#[derive(Debug, Clone, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct AXProperty {
pub name: String,
pub value: AXValue,
}
// ---------------------------------------------------------------------------
// Network domain (minimal for Phase 1)
// ---------------------------------------------------------------------------
#[derive(Debug, Clone, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct RequestWillBeSentEvent {
pub request_id: String,
pub request: NetworkRequest,
}
#[derive(Debug, Clone, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct NetworkRequest {
pub url: String,
pub method: String,
}
#[derive(Debug, Clone, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct LoadingFinishedEvent {
pub request_id: String,
}
#[derive(Debug, Clone, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct LoadingFailedEvent {
pub request_id: String,
}
// ---------------------------------------------------------------------------
// DOM domain
// ---------------------------------------------------------------------------
#[derive(Debug, Serialize)]
#[serde(rename_all = "camelCase")]
pub struct DomResolveNodeParams {
#[serde(skip_serializing_if = "Option::is_none")]
pub backend_node_id: Option<i64>,
#[serde(skip_serializing_if = "Option::is_none")]
pub node_id: Option<i64>,
#[serde(skip_serializing_if = "Option::is_none")]
pub object_group: Option<String>,
}
#[derive(Debug, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct DomResolveNodeResult {
pub object: RemoteObject,
}
#[derive(Debug, Serialize)]
#[serde(rename_all = "camelCase")]
pub struct DomGetBoxModelParams {
#[serde(skip_serializing_if = "Option::is_none")]
pub backend_node_id: Option<i64>,
#[serde(skip_serializing_if = "Option::is_none")]
pub node_id: Option<i64>,
#[serde(skip_serializing_if = "Option::is_none")]
pub object_id: Option<String>,
}
#[derive(Debug, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct DomGetBoxModelResult {
pub model: BoxModel,
}
#[derive(Debug, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct BoxModel {
pub content: Vec<f64>,
pub padding: Vec<f64>,
pub border: Vec<f64>,
pub margin: Vec<f64>,
pub width: i64,
pub height: i64,
}
#[derive(Debug, Serialize)]
#[serde(rename_all = "camelCase")]
pub struct DomQuerySelectorParams {
pub node_id: i64,
pub selector: String,
}
#[derive(Debug, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct DomQuerySelectorResult {
pub node_id: i64,
}
#[derive(Debug, Serialize)]
#[serde(rename_all = "camelCase")]
pub struct DomGetDocumentParams {
#[serde(skip_serializing_if = "Option::is_none")]
pub depth: Option<i32>,
}
#[derive(Debug, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct DomGetDocumentResult {
pub root: DomNode,
}
#[derive(Debug, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct DomNode {
pub node_id: i64,
pub backend_node_id: Option<i64>,
pub node_type: Option<i64>,
pub node_name: Option<String>,
pub children: Option<Vec<DomNode>>,
}
// ---------------------------------------------------------------------------
// Input domain
// ---------------------------------------------------------------------------
#[derive(Debug, Serialize)]
#[serde(rename_all = "camelCase")]
pub struct DispatchMouseEventParams {
#[serde(rename = "type")]
pub event_type: String,
pub x: f64,
pub y: f64,
#[serde(skip_serializing_if = "Option::is_none")]
pub button: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub buttons: Option<i32>,
#[serde(skip_serializing_if = "Option::is_none")]
pub click_count: Option<i32>,
#[serde(skip_serializing_if = "Option::is_none")]
pub delta_x: Option<f64>,
#[serde(skip_serializing_if = "Option::is_none")]
pub delta_y: Option<f64>,
#[serde(skip_serializing_if = "Option::is_none")]
pub modifiers: Option<i32>,
}
#[derive(Debug, Serialize)]
#[serde(rename_all = "camelCase")]
pub struct DispatchKeyEventParams {
#[serde(rename = "type")]
pub event_type: String,
#[serde(skip_serializing_if = "Option::is_none")]
pub key: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub code: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub text: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub unmodified_text: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub windows_virtual_key_code: Option<i32>,
#[serde(skip_serializing_if = "Option::is_none")]
pub native_virtual_key_code: Option<i32>,
#[serde(skip_serializing_if = "Option::is_none")]
pub modifiers: Option<i32>,
}
#[derive(Debug, Serialize)]
#[serde(rename_all = "camelCase")]
pub struct InsertTextParams {
pub text: String,
}
// ---------------------------------------------------------------------------
// Page.captureScreenshot
// ---------------------------------------------------------------------------
#[derive(Debug, Serialize)]
#[serde(rename_all = "camelCase")]
pub struct CaptureScreenshotParams {
#[serde(skip_serializing_if = "Option::is_none")]
pub format: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub quality: Option<i32>,
#[serde(skip_serializing_if = "Option::is_none")]
pub clip: Option<Viewport>,
#[serde(skip_serializing_if = "Option::is_none")]
pub from_surface: Option<bool>,
#[serde(skip_serializing_if = "Option::is_none")]
pub capture_beyond_viewport: Option<bool>,
}
#[derive(Debug, Serialize)]
#[serde(rename_all = "camelCase")]
pub struct Viewport {
pub x: f64,
pub y: f64,
pub width: f64,
pub height: f64,
pub scale: f64,
}
#[derive(Debug, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct CaptureScreenshotResult {
pub data: String,
}
// ---------------------------------------------------------------------------
// Runtime.callFunctionOn
// ---------------------------------------------------------------------------
#[derive(Debug, Serialize)]
#[serde(rename_all = "camelCase")]
pub struct CallFunctionOnParams {
pub function_declaration: String,
#[serde(skip_serializing_if = "Option::is_none")]
pub object_id: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub arguments: Option<Vec<CallArgument>>,
#[serde(skip_serializing_if = "Option::is_none")]
pub return_by_value: Option<bool>,
#[serde(skip_serializing_if = "Option::is_none")]
pub await_promise: Option<bool>,
}
#[derive(Debug, Serialize)]
#[serde(rename_all = "camelCase")]
pub struct CallArgument {
#[serde(skip_serializing_if = "Option::is_none")]
pub value: Option<Value>,
#[serde(skip_serializing_if = "Option::is_none")]
pub object_id: Option<String>,
}
// ---------------------------------------------------------------------------
// Version info (from /json/version)
// ---------------------------------------------------------------------------
#[derive(Debug, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct BrowserVersionInfo {
#[serde(rename = "webSocketDebuggerUrl")]
pub web_socket_debugger_url: Option<String>,
#[serde(rename = "Browser")]
pub browser: Option<String>,
}
/// Auto-generated CDP types from protocol JSON files in `cdp-protocol/`.
///
/// To populate: download `browser_protocol.json` and `js_protocol.json` from
/// <https://github.com/nicolo-ribaudo/nicolo-ribaudo.github.io/> (or any
/// Chromium source) into `cli/cdp-protocol/` and rebuild.
///
/// Usage: `use super::cdp::types::generated::cdp_page::*;`
pub mod generated {
include!(concat!(env!("OUT_DIR"), "/cdp_generated.rs"));
}
+87
View File
@@ -0,0 +1,87 @@
use serde::{Deserialize, Serialize};
use serde_json::{json, Value};
use super::cdp::client::CdpClient;
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct Cookie {
pub name: String,
pub value: String,
pub domain: String,
pub path: String,
#[serde(default)]
pub expires: f64,
#[serde(default)]
pub size: i64,
#[serde(default)]
pub http_only: bool,
#[serde(default)]
pub secure: bool,
#[serde(default)]
pub session: bool,
#[serde(skip_serializing_if = "Option::is_none")]
pub same_site: Option<String>,
}
pub async fn get_cookies(
client: &CdpClient,
session_id: &str,
urls: Option<Vec<String>>,
) -> Result<Vec<Cookie>, String> {
let params = match urls {
Some(ref u) if !u.is_empty() => json!({ "urls": u }),
_ => json!({}),
};
let result = client
.send_command("Network.getCookies", Some(params), Some(session_id))
.await?;
let cookies: Vec<Cookie> = result
.get("cookies")
.and_then(|v| serde_json::from_value(v.clone()).ok())
.unwrap_or_default();
Ok(cookies)
}
pub async fn set_cookies(
client: &CdpClient,
session_id: &str,
cookies: Vec<Value>,
current_url: Option<&str>,
) -> Result<(), String> {
let cookies: Vec<Value> = cookies
.into_iter()
.map(|mut c| {
// Auto-fill url if no domain/path/url provided
if c.get("url").is_none() && c.get("domain").is_none() && current_url.is_some() {
c.as_object_mut().map(|m| {
m.insert(
"url".to_string(),
Value::String(current_url.unwrap().to_string()),
)
});
}
c
})
.collect();
client
.send_command(
"Network.setCookies",
Some(json!({ "cookies": cookies })),
Some(session_id),
)
.await?;
Ok(())
}
pub async fn clear_cookies(client: &CdpClient, session_id: &str) -> Result<(), String> {
client
.send_command_no_params("Network.clearBrowserCookies", Some(session_id))
.await?;
Ok(())
}
+266
View File
@@ -0,0 +1,266 @@
use serde_json::Value;
use std::env;
use std::fs;
use std::path::PathBuf;
use std::process;
use tokio::io::{AsyncBufReadExt, AsyncWriteExt, BufReader};
use tokio::signal;
use super::actions::{execute_command, DaemonState};
use super::state;
pub async fn run_daemon(session: &str) {
let socket_dir = get_daemon_socket_dir();
if !socket_dir.exists() {
let _ = fs::create_dir_all(&socket_dir);
}
let pid_path = socket_dir.join(format!("{}.pid", session));
let _ = fs::write(&pid_path, process::id().to_string());
let socket_path = socket_dir.join(format!("{}.sock", session));
if socket_path.exists() {
let _ = fs::remove_file(&socket_path);
}
if let Ok(days_str) = env::var("AGENT_BROWSER_STATE_EXPIRE_DAYS") {
if let Ok(days) = days_str.parse::<u64>() {
if days > 0 {
let _ = state::state_clean(days);
}
}
}
let result = run_socket_server(&socket_path, session).await;
let _ = fs::remove_file(&socket_path);
let _ = fs::remove_file(&pid_path);
let stream_path = socket_dir.join(format!("{}.stream", session));
let _ = fs::remove_file(&stream_path);
if let Err(e) = result {
eprintln!("Daemon error: {}", e);
process::exit(1);
}
}
#[cfg(unix)]
async fn run_socket_server(socket_path: &PathBuf, _session: &str) -> Result<(), String> {
use tokio::net::UnixListener;
let listener =
UnixListener::bind(socket_path).map_err(|e| format!("Failed to bind socket: {}", e))?;
let state: std::sync::Arc<tokio::sync::Mutex<DaemonState>> =
std::sync::Arc::new(tokio::sync::Mutex::new(DaemonState::new()));
loop {
tokio::select! {
accept_result = listener.accept() => {
match accept_result {
Ok((stream, _)) => {
let state = state.clone();
tokio::spawn(async move {
handle_connection(stream, state).await;
});
}
Err(e) => {
eprintln!("Accept error: {}", e);
}
}
}
_ = shutdown_signal() => {
let mut s = state.lock().await;
if let Some(ref mut mgr) = s.browser {
let _ = mgr.close().await;
}
break;
}
}
}
Ok(())
}
#[cfg(windows)]
async fn run_socket_server(socket_path: &PathBuf, session: &str) -> Result<(), String> {
use tokio::net::TcpListener;
let port = get_port_for_session(session);
let listener = TcpListener::bind(format!("127.0.0.1:{}", port))
.await
.map_err(|e| format!("Failed to bind TCP: {}", e))?;
let socket_dir = socket_path.parent().unwrap_or(std::path::Path::new("."));
let port_path = socket_dir.join(format!("{}.port", session));
let _ = fs::write(&port_path, port.to_string());
let state: std::sync::Arc<tokio::sync::Mutex<DaemonState>> =
std::sync::Arc::new(tokio::sync::Mutex::new(DaemonState::new()));
loop {
tokio::select! {
accept_result = listener.accept() => {
match accept_result {
Ok((stream, _)) => {
let state = state.clone();
tokio::spawn(async move {
handle_connection(stream, state).await;
});
}
Err(e) => {
eprintln!("Accept error: {}", e);
}
}
}
_ = shutdown_signal() => {
let mut s = state.lock().await;
if let Some(ref mut mgr) = s.browser {
let _ = mgr.close().await;
}
let _ = fs::remove_file(&port_path);
break;
}
}
}
Ok(())
}
async fn handle_connection<S>(stream: S, state: std::sync::Arc<tokio::sync::Mutex<DaemonState>>)
where
S: tokio::io::AsyncRead + tokio::io::AsyncWrite + Unpin,
{
let (reader, mut writer) = tokio::io::split(stream);
let mut buf_reader = BufReader::new(reader);
let mut line = String::new();
loop {
line.clear();
match buf_reader.read_line(&mut line).await {
Ok(0) => break,
Ok(_) => {
let trimmed = line.trim();
if trimmed.is_empty() {
continue;
}
if looks_like_http(trimmed) {
break;
}
let cmd: Value = match serde_json::from_str(trimmed) {
Ok(v) => v,
Err(e) => {
let err = serde_json::json!({
"success": false,
"error": format!("Invalid JSON: {}", e),
});
let mut resp = serde_json::to_string(&err).unwrap_or_default();
resp.push('\n');
let _ = writer.write_all(resp.as_bytes()).await;
continue;
}
};
let is_close = cmd.get("action").and_then(|v| v.as_str()) == Some("close");
let response = {
let mut s = state.lock().await;
execute_command(&cmd, &mut s).await
};
let mut resp = serde_json::to_string(&response).unwrap_or_default();
resp.push('\n');
if writer.write_all(resp.as_bytes()).await.is_err() {
break;
}
if is_close {
tokio::time::sleep(tokio::time::Duration::from_millis(100)).await;
process::exit(0);
}
}
Err(_) => break,
}
}
}
fn looks_like_http(line: &str) -> bool {
let prefixes = [
"GET ", "POST ", "PUT ", "DELETE ", "PATCH ", "HEAD ", "OPTIONS ", "CONNECT ", "TRACE ",
];
prefixes.iter().any(|p| line.starts_with(p))
}
async fn shutdown_signal() {
#[cfg(unix)]
{
let mut sigint = match signal::unix::signal(signal::unix::SignalKind::interrupt()) {
Ok(s) => s,
Err(e) => {
eprintln!("Failed to install SIGINT handler: {}", e);
process::exit(1);
}
};
let mut sigterm = match signal::unix::signal(signal::unix::SignalKind::terminate()) {
Ok(s) => s,
Err(e) => {
eprintln!("Failed to install SIGTERM handler: {}", e);
process::exit(1);
}
};
let mut sighup = match signal::unix::signal(signal::unix::SignalKind::hangup()) {
Ok(s) => s,
Err(e) => {
eprintln!("Failed to install SIGHUP handler: {}", e);
process::exit(1);
}
};
tokio::select! {
_ = sigint.recv() => {}
_ = sigterm.recv() => {}
_ = sighup.recv() => {}
}
}
#[cfg(windows)]
{
if let Err(e) = signal::ctrl_c().await {
eprintln!("Failed to install Ctrl+C handler: {}", e);
process::exit(1);
}
}
}
fn get_daemon_socket_dir() -> PathBuf {
if let Ok(dir) = env::var("AGENT_BROWSER_SOCKET_DIR") {
if !dir.is_empty() {
return PathBuf::from(dir);
}
}
if let Ok(xdg) = env::var("XDG_RUNTIME_DIR") {
if !xdg.is_empty() {
return PathBuf::from(xdg).join("agent-browser");
}
}
if let Some(home) = dirs::home_dir() {
return home.join(".agent-browser");
}
std::env::temp_dir().join("agent-browser")
}
#[cfg(windows)]
fn get_port_for_session(session: &str) -> u16 {
let mut hash: i64 = 0;
for b in session.bytes() {
hash = hash.wrapping_mul(31).wrapping_add(b as i64);
}
49152 + (hash.unsigned_abs() % 16383) as u16
}
+195
View File
@@ -0,0 +1,195 @@
use serde_json::{json, Value};
use similar::{ChangeTag, TextDiff};
pub struct ScreenshotDiffResult {
pub total_pixels: u64,
pub different_pixels: u64,
pub mismatch_percentage: f64,
pub matched: bool,
pub diff_image: Option<Vec<u8>>,
pub dimension_mismatch: Option<Value>,
}
pub struct SnapshotDiffResult {
pub diff: String,
pub additions: usize,
pub removals: usize,
pub unchanged: usize,
pub changed: bool,
}
pub fn diff_screenshot(
baseline: &[u8],
current: &[u8],
threshold: f64,
) -> Result<ScreenshotDiffResult, String> {
let img_a = image::load_from_memory(baseline)
.map_err(|e| format!("Failed to decode baseline image: {}", e))?;
let img_b = image::load_from_memory(current)
.map_err(|e| format!("Failed to decode current image: {}", e))?;
let (wa, ha) = (img_a.width(), img_a.height());
let (wb, hb) = (img_b.width(), img_b.height());
if wa != wb || ha != hb {
return Ok(ScreenshotDiffResult {
total_pixels: (wa as u64) * (ha as u64),
different_pixels: (wa as u64) * (ha as u64),
mismatch_percentage: 100.0,
matched: false,
diff_image: None,
dimension_mismatch: Some(json!({
"expected": { "width": wa, "height": ha },
"actual": { "width": wb, "height": hb },
})),
});
}
let rgba_a = img_a.to_rgba8();
let rgba_b = img_b.to_rgba8();
let total = (wa as u64) * (ha as u64);
let max_color_distance = threshold * 255.0 * (3.0_f64).sqrt();
let mut different = 0u64;
let mut diff_img = image::RgbaImage::new(wa, ha);
for y in 0..ha {
for x in 0..wa {
let pa = rgba_a.get_pixel(x, y);
let pb = rgba_b.get_pixel(x, y);
let dr = (pa[0] as f64) - (pb[0] as f64);
let dg = (pa[1] as f64) - (pb[1] as f64);
let db = (pa[2] as f64) - (pb[2] as f64);
let dist = (dr * dr + dg * dg + db * db).sqrt();
if dist > max_color_distance {
different += 1;
diff_img.put_pixel(x, y, image::Rgba([255, 0, 0, 255]));
} else {
let gray = ((pa[0] as u16 + pa[1] as u16 + pa[2] as u16) / 3) as u8;
let dimmed = (gray as f64 * 0.3) as u8;
diff_img.put_pixel(x, y, image::Rgba([dimmed, dimmed, dimmed, 255]));
}
}
}
let mismatch = if total > 0 {
(different as f64 / total as f64) * 100.0
} else {
0.0
};
let diff_bytes = if different > 0 {
let mut buf = std::io::Cursor::new(Vec::new());
diff_img
.write_to(&mut buf, image::ImageFormat::Png)
.map_err(|e| format!("Failed to encode diff image: {}", e))?;
Some(buf.into_inner())
} else {
None
};
Ok(ScreenshotDiffResult {
total_pixels: total,
different_pixels: different,
mismatch_percentage: mismatch,
matched: different == 0,
diff_image: diff_bytes,
dimension_mismatch: None,
})
}
/// Compute a snapshot diff using the Myers algorithm via the `similar` crate.
pub fn diff_snapshots(before: &str, after: &str) -> SnapshotDiffResult {
let text_diff = TextDiff::from_lines(before, after);
let mut additions = 0usize;
let mut removals = 0usize;
let mut unchanged = 0usize;
for change in text_diff.iter_all_changes() {
match change.tag() {
ChangeTag::Insert => additions += 1,
ChangeTag::Delete => removals += 1,
ChangeTag::Equal => unchanged += 1,
}
}
let changed = additions > 0 || removals > 0;
let diff = text_diff
.unified_diff()
.context_radius(3)
.header("before", "after")
.to_string();
SnapshotDiffResult {
diff,
additions,
removals,
unchanged,
changed,
}
}
/// Legacy JSON diff output for backwards compatibility.
pub fn diff_text(a: &str, b: &str) -> Value {
let result = diff_snapshots(a, b);
json!({
"identical": !result.changed,
"additions": result.additions,
"removals": result.removals,
"deletions": result.removals,
"unchanged": result.unchanged,
"changed": result.changed,
})
}
pub fn diff_unified(a: &str, b: &str) -> String {
diff_snapshots(a, b).diff
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_diff_identical() {
let result = diff_text("hello\nworld", "hello\nworld");
assert_eq!(result.get("identical").unwrap(), true);
assert_eq!(result.get("changed").unwrap(), false);
assert_eq!(result.get("unchanged").unwrap(), 2);
}
#[test]
fn test_diff_additions() {
let result = diff_text("hello\n", "hello\nworld\n");
assert_eq!(result.get("identical").unwrap(), false);
assert_eq!(result.get("changed").unwrap(), true);
assert!(result.get("additions").unwrap().as_i64().unwrap() > 0);
}
#[test]
fn test_diff_deletions() {
let result = diff_text("hello\nworld\n", "hello\n");
assert_eq!(result.get("identical").unwrap(), false);
assert!(result.get("removals").unwrap().as_i64().unwrap() > 0);
}
#[test]
fn test_diff_unified_output() {
let output = diff_unified("a\nb\nc\n", "a\nx\nc\n");
assert!(output.contains("---"));
assert!(output.contains("+++"));
}
#[test]
fn test_snapshot_diff_struct() {
let result = diff_snapshots("line1\nline2\n", "line1\nline3\n");
assert!(result.changed);
assert_eq!(result.additions, 1);
assert_eq!(result.removals, 1);
assert_eq!(result.unchanged, 1);
assert!(!result.diff.is_empty());
}
}
File diff suppressed because it is too large Load Diff
+718
View File
@@ -0,0 +1,718 @@
use std::collections::HashMap;
use serde_json::Value;
use super::cdp::client::CdpClient;
use super::cdp::types::*;
#[derive(Debug, Clone)]
pub struct RefEntry {
pub backend_node_id: Option<i64>,
pub role: String,
pub name: String,
pub nth: Option<usize>,
pub selector: Option<String>,
}
pub struct RefMap {
map: HashMap<String, RefEntry>,
next_ref: usize,
}
impl RefMap {
pub fn new() -> Self {
Self {
map: HashMap::new(),
next_ref: 1,
}
}
pub fn add(
&mut self,
ref_id: String,
backend_node_id: Option<i64>,
role: &str,
name: &str,
nth: Option<usize>,
) {
self.map.insert(
ref_id,
RefEntry {
backend_node_id,
role: role.to_string(),
name: name.to_string(),
nth,
selector: None,
},
);
}
pub fn get(&self, ref_id: &str) -> Option<&RefEntry> {
self.map.get(ref_id)
}
pub fn clear(&mut self) {
self.map.clear();
self.next_ref = 1;
}
pub fn next_ref_num(&self) -> usize {
self.next_ref
}
pub fn set_next_ref_num(&mut self, n: usize) {
self.next_ref = n;
}
}
pub fn parse_ref(input: &str) -> Option<String> {
let trimmed = input.trim();
if let Some(stripped) = trimmed.strip_prefix('@') {
if stripped.starts_with('e') && stripped[1..].chars().all(|c| c.is_ascii_digit()) {
return Some(stripped.to_string());
}
}
if let Some(stripped) = trimmed.strip_prefix("ref=") {
if stripped.starts_with('e') && stripped[1..].chars().all(|c| c.is_ascii_digit()) {
return Some(stripped.to_string());
}
}
if trimmed.starts_with('e')
&& trimmed.len() > 1
&& trimmed[1..].chars().all(|c| c.is_ascii_digit())
{
return Some(trimmed.to_string());
}
None
}
pub async fn resolve_element_center(
client: &CdpClient,
session_id: &str,
ref_map: &RefMap,
selector_or_ref: &str,
) -> Result<(f64, f64), String> {
if let Some(ref_id) = parse_ref(selector_or_ref) {
let entry = ref_map
.get(&ref_id)
.ok_or_else(|| format!("Unknown ref: {}", ref_id))?;
if let Some(backend_node_id) = entry.backend_node_id {
let result: DomGetBoxModelResult = client
.send_command_typed(
"DOM.getBoxModel",
&DomGetBoxModelParams {
backend_node_id: Some(backend_node_id),
node_id: None,
object_id: None,
},
Some(session_id),
)
.await?;
return Ok(box_model_center(&result.model));
}
// Fallback: use role/name to find via JS
return resolve_by_role_name(client, session_id, &entry.role, &entry.name, entry.nth).await;
}
// CSS selector
resolve_by_selector(client, session_id, selector_or_ref).await
}
pub async fn resolve_element_object_id(
client: &CdpClient,
session_id: &str,
ref_map: &RefMap,
selector_or_ref: &str,
) -> Result<String, String> {
if let Some(ref_id) = parse_ref(selector_or_ref) {
let entry = ref_map
.get(&ref_id)
.ok_or_else(|| format!("Unknown ref: {}", ref_id))?;
if let Some(backend_node_id) = entry.backend_node_id {
let result: DomResolveNodeResult = client
.send_command_typed(
"DOM.resolveNode",
&DomResolveNodeParams {
backend_node_id: Some(backend_node_id),
node_id: None,
object_group: Some("agent-browser".to_string()),
},
Some(session_id),
)
.await?;
return result
.object
.object_id
.ok_or_else(|| format!("No objectId for ref {}", ref_id));
}
}
// CSS selector fallback
let js = format!(
"document.querySelector({})",
serde_json::to_string(selector_or_ref).unwrap_or_default()
);
let result: EvaluateResult = client
.send_command_typed(
"Runtime.evaluate",
&EvaluateParams {
expression: js,
return_by_value: Some(false),
await_promise: Some(false),
},
Some(session_id),
)
.await?;
result
.result
.object_id
.ok_or_else(|| format!("Element not found: {}", selector_or_ref))
}
async fn resolve_by_role_name(
client: &CdpClient,
session_id: &str,
role: &str,
name: &str,
nth: Option<usize>,
) -> Result<(f64, f64), String> {
let nth_index = nth.unwrap_or(0);
let js = format!(
r#"(() => {{
const walker = document.createTreeWalker(document.body, NodeFilter.SHOW_ELEMENT);
const matches = [];
let node;
while (node = walker.nextNode()) {{
const r = node.getAttribute('role') || node.tagName.toLowerCase();
const n = node.getAttribute('aria-label') || node.textContent.trim().slice(0, 100);
if (r === {role} && n === {name}) matches.push(node);
}}
const el = matches[{nth}];
if (!el) return null;
const rect = el.getBoundingClientRect();
return {{ x: rect.x + rect.width / 2, y: rect.y + rect.height / 2 }};
}})()"#,
role = serde_json::to_string(role).unwrap_or_default(),
name = serde_json::to_string(name).unwrap_or_default(),
nth = nth_index,
);
let result: EvaluateResult = client
.send_command_typed(
"Runtime.evaluate",
&EvaluateParams {
expression: js,
return_by_value: Some(true),
await_promise: Some(false),
},
Some(session_id),
)
.await?;
let val = result.result.value.unwrap_or(Value::Null);
let x = val.get("x").and_then(|v| v.as_f64());
let y = val.get("y").and_then(|v| v.as_f64());
match (x, y) {
(Some(x), Some(y)) => Ok((x, y)),
_ => Err(format!(
"Could not locate element with role={} name={}",
role, name
)),
}
}
async fn resolve_by_selector(
client: &CdpClient,
session_id: &str,
selector: &str,
) -> Result<(f64, f64), String> {
let js = format!(
r#"(() => {{
const el = document.querySelector({sel});
if (!el) return null;
const rect = el.getBoundingClientRect();
return {{ x: rect.x + rect.width / 2, y: rect.y + rect.height / 2 }};
}})()"#,
sel = serde_json::to_string(selector).unwrap_or_default(),
);
let result: EvaluateResult = client
.send_command_typed(
"Runtime.evaluate",
&EvaluateParams {
expression: js,
return_by_value: Some(true),
await_promise: Some(false),
},
Some(session_id),
)
.await?;
let val = result.result.value.unwrap_or(Value::Null);
let x = val.get("x").and_then(|v| v.as_f64());
let y = val.get("y").and_then(|v| v.as_f64());
match (x, y) {
(Some(x), Some(y)) => Ok((x, y)),
_ => Err(format!("Element not found: {}", selector)),
}
}
fn box_model_center(model: &BoxModel) -> (f64, f64) {
// content quad: [x1,y1, x2,y2, x3,y3, x4,y4]
if model.content.len() >= 8 {
let x = (model.content[0] + model.content[2] + model.content[4] + model.content[6]) / 4.0;
let y = (model.content[1] + model.content[3] + model.content[5] + model.content[7]) / 4.0;
(x, y)
} else {
(0.0, 0.0)
}
}
pub async fn get_element_text(
client: &CdpClient,
session_id: &str,
ref_map: &RefMap,
selector_or_ref: &str,
) -> Result<String, String> {
let object_id = resolve_element_object_id(client, session_id, ref_map, selector_or_ref).await?;
let result: EvaluateResult = client
.send_command_typed(
"Runtime.callFunctionOn",
&CallFunctionOnParams {
function_declaration:
"function() { return this.innerText || this.textContent || ''; }".to_string(),
object_id: Some(object_id),
arguments: None,
return_by_value: Some(true),
await_promise: Some(false),
},
Some(session_id),
)
.await?;
Ok(result
.result
.value
.and_then(|v| v.as_str().map(|s| s.to_string()))
.unwrap_or_default())
}
pub async fn get_element_attribute(
client: &CdpClient,
session_id: &str,
ref_map: &RefMap,
selector_or_ref: &str,
attribute: &str,
) -> Result<Value, String> {
let object_id = resolve_element_object_id(client, session_id, ref_map, selector_or_ref).await?;
let result: EvaluateResult = client
.send_command_typed(
"Runtime.callFunctionOn",
&CallFunctionOnParams {
function_declaration: format!(
"function() {{ return this.getAttribute({}); }}",
serde_json::to_string(attribute).unwrap_or_default()
),
object_id: Some(object_id),
arguments: None,
return_by_value: Some(true),
await_promise: Some(false),
},
Some(session_id),
)
.await?;
Ok(result.result.value.unwrap_or(Value::Null))
}
pub async fn is_element_visible(
client: &CdpClient,
session_id: &str,
ref_map: &RefMap,
selector_or_ref: &str,
) -> Result<bool, String> {
let object_id = resolve_element_object_id(client, session_id, ref_map, selector_or_ref).await?;
let result: EvaluateResult = client
.send_command_typed(
"Runtime.callFunctionOn",
&CallFunctionOnParams {
function_declaration: r#"function() {
const rect = this.getBoundingClientRect();
const style = window.getComputedStyle(this);
return rect.width > 0 && rect.height > 0 &&
style.visibility !== 'hidden' &&
style.display !== 'none' &&
parseFloat(style.opacity) > 0;
}"#
.to_string(),
object_id: Some(object_id),
arguments: None,
return_by_value: Some(true),
await_promise: Some(false),
},
Some(session_id),
)
.await?;
Ok(result
.result
.value
.and_then(|v| v.as_bool())
.unwrap_or(false))
}
pub async fn is_element_enabled(
client: &CdpClient,
session_id: &str,
ref_map: &RefMap,
selector_or_ref: &str,
) -> Result<bool, String> {
let object_id = resolve_element_object_id(client, session_id, ref_map, selector_or_ref).await?;
let result: EvaluateResult = client
.send_command_typed(
"Runtime.callFunctionOn",
&CallFunctionOnParams {
function_declaration: "function() { return !this.disabled; }".to_string(),
object_id: Some(object_id),
arguments: None,
return_by_value: Some(true),
await_promise: Some(false),
},
Some(session_id),
)
.await?;
Ok(result
.result
.value
.and_then(|v| v.as_bool())
.unwrap_or(true))
}
pub async fn is_element_checked(
client: &CdpClient,
session_id: &str,
ref_map: &RefMap,
selector_or_ref: &str,
) -> Result<bool, String> {
let object_id = resolve_element_object_id(client, session_id, ref_map, selector_or_ref).await?;
let result: EvaluateResult = client
.send_command_typed(
"Runtime.callFunctionOn",
&CallFunctionOnParams {
function_declaration: "function() { return !!this.checked; }".to_string(),
object_id: Some(object_id),
arguments: None,
return_by_value: Some(true),
await_promise: Some(false),
},
Some(session_id),
)
.await?;
Ok(result
.result
.value
.and_then(|v| v.as_bool())
.unwrap_or(false))
}
pub async fn get_element_inner_text(
client: &CdpClient,
session_id: &str,
ref_map: &RefMap,
selector_or_ref: &str,
) -> Result<String, String> {
let object_id = resolve_element_object_id(client, session_id, ref_map, selector_or_ref).await?;
let result: EvaluateResult = client
.send_command_typed(
"Runtime.callFunctionOn",
&CallFunctionOnParams {
function_declaration: "function() { return this.innerText || ''; }".to_string(),
object_id: Some(object_id),
arguments: None,
return_by_value: Some(true),
await_promise: Some(false),
},
Some(session_id),
)
.await?;
Ok(result
.result
.value
.and_then(|v| v.as_str().map(|s| s.to_string()))
.unwrap_or_default())
}
pub async fn get_element_inner_html(
client: &CdpClient,
session_id: &str,
ref_map: &RefMap,
selector_or_ref: &str,
) -> Result<String, String> {
let object_id = resolve_element_object_id(client, session_id, ref_map, selector_or_ref).await?;
let result: EvaluateResult = client
.send_command_typed(
"Runtime.callFunctionOn",
&CallFunctionOnParams {
function_declaration: "function() { return this.innerHTML || ''; }".to_string(),
object_id: Some(object_id),
arguments: None,
return_by_value: Some(true),
await_promise: Some(false),
},
Some(session_id),
)
.await?;
Ok(result
.result
.value
.and_then(|v| v.as_str().map(|s| s.to_string()))
.unwrap_or_default())
}
pub async fn get_element_input_value(
client: &CdpClient,
session_id: &str,
ref_map: &RefMap,
selector_or_ref: &str,
) -> Result<String, String> {
let object_id = resolve_element_object_id(client, session_id, ref_map, selector_or_ref).await?;
let result: EvaluateResult = client
.send_command_typed(
"Runtime.callFunctionOn",
&CallFunctionOnParams {
function_declaration:
"function() { return typeof this.value === 'string' ? this.value : ''; }"
.to_string(),
object_id: Some(object_id),
arguments: None,
return_by_value: Some(true),
await_promise: Some(false),
},
Some(session_id),
)
.await?;
Ok(result
.result
.value
.and_then(|v| v.as_str().map(|s| s.to_string()))
.unwrap_or_default())
}
pub async fn set_element_value(
client: &CdpClient,
session_id: &str,
ref_map: &RefMap,
selector_or_ref: &str,
value: &str,
) -> Result<(), String> {
let object_id = resolve_element_object_id(client, session_id, ref_map, selector_or_ref).await?;
let js = format!(
"function() {{ this.value = {}; this.dispatchEvent(new Event('input', {{bubbles: true}})); this.dispatchEvent(new Event('change', {{bubbles: true}})); }}",
serde_json::to_string(value).unwrap_or_default()
);
client
.send_command_typed::<_, EvaluateResult>(
"Runtime.callFunctionOn",
&CallFunctionOnParams {
function_declaration: js,
object_id: Some(object_id),
arguments: None,
return_by_value: Some(true),
await_promise: Some(false),
},
Some(session_id),
)
.await?;
Ok(())
}
pub async fn get_element_bounding_box(
client: &CdpClient,
session_id: &str,
ref_map: &RefMap,
selector_or_ref: &str,
) -> Result<Value, String> {
let object_id = resolve_element_object_id(client, session_id, ref_map, selector_or_ref).await?;
let result: EvaluateResult = client
.send_command_typed(
"Runtime.callFunctionOn",
&CallFunctionOnParams {
function_declaration: r#"function() {
const r = this.getBoundingClientRect();
return { x: r.x, y: r.y, width: r.width, height: r.height };
}"#
.to_string(),
object_id: Some(object_id),
arguments: None,
return_by_value: Some(true),
await_promise: Some(false),
},
Some(session_id),
)
.await?;
result
.result
.value
.ok_or_else(|| format!("Could not get bounding box for: {}", selector_or_ref))
}
pub async fn get_element_count(
client: &CdpClient,
session_id: &str,
selector: &str,
) -> Result<i64, String> {
let js = format!(
"document.querySelectorAll({}).length",
serde_json::to_string(selector).unwrap_or_default()
);
let result: EvaluateResult = client
.send_command_typed(
"Runtime.evaluate",
&EvaluateParams {
expression: js,
return_by_value: Some(true),
await_promise: Some(false),
},
Some(session_id),
)
.await?;
Ok(result.result.value.and_then(|v| v.as_i64()).unwrap_or(0))
}
pub async fn get_element_styles(
client: &CdpClient,
session_id: &str,
ref_map: &RefMap,
selector_or_ref: &str,
properties: Option<Vec<String>>,
) -> Result<Value, String> {
let object_id = resolve_element_object_id(client, session_id, ref_map, selector_or_ref).await?;
let js = match properties {
Some(props) => {
let props_json = serde_json::to_string(&props).unwrap_or("[]".to_string());
format!(
r#"function() {{
const s = window.getComputedStyle(this);
const props = {};
const result = {{}};
for (const p of props) result[p] = s.getPropertyValue(p);
return result;
}}"#,
props_json
)
}
None => r#"function() {
const s = window.getComputedStyle(this);
const result = {};
for (let i = 0; i < s.length; i++) {
const p = s[i];
result[p] = s.getPropertyValue(p);
}
return result;
}"#
.to_string(),
};
let result: EvaluateResult = client
.send_command_typed(
"Runtime.callFunctionOn",
&CallFunctionOnParams {
function_declaration: js,
object_id: Some(object_id),
arguments: None,
return_by_value: Some(true),
await_promise: Some(false),
},
Some(session_id),
)
.await?;
Ok(result.result.value.unwrap_or(Value::Null))
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_parse_ref_at_prefix() {
assert_eq!(parse_ref("@e1"), Some("e1".to_string()));
assert_eq!(parse_ref("@e123"), Some("e123".to_string()));
}
#[test]
fn test_parse_ref_equals_prefix() {
assert_eq!(parse_ref("ref=e1"), Some("e1".to_string()));
}
#[test]
fn test_parse_ref_bare() {
assert_eq!(parse_ref("e1"), Some("e1".to_string()));
assert_eq!(parse_ref("e42"), Some("e42".to_string()));
}
#[test]
fn test_parse_ref_invalid() {
assert_eq!(parse_ref("button"), None);
assert_eq!(parse_ref("e"), None);
assert_eq!(parse_ref("1"), None);
assert_eq!(parse_ref(""), None);
}
#[test]
fn test_ref_map_basic() {
let mut map = RefMap::new();
map.add("e1".to_string(), Some(42), "button", "Submit", None);
assert!(map.get("e1").is_some());
assert_eq!(map.get("e1").unwrap().role, "button");
assert!(map.get("e2").is_none());
}
#[test]
fn test_box_model_center() {
let model = BoxModel {
content: vec![10.0, 20.0, 110.0, 20.0, 110.0, 60.0, 10.0, 60.0],
padding: vec![],
border: vec![],
margin: vec![],
width: 100,
height: 40,
};
let (x, y) = box_model_center(&model);
assert!((x - 60.0).abs() < 0.01);
assert!((y - 40.0).abs() < 0.01);
}
}
+707
View File
@@ -0,0 +1,707 @@
use serde_json::Value;
use super::cdp::client::CdpClient;
use super::cdp::types::*;
use super::element::{resolve_element_center, resolve_element_object_id, RefMap};
pub async fn click(
client: &CdpClient,
session_id: &str,
ref_map: &RefMap,
selector_or_ref: &str,
button: &str,
click_count: i32,
) -> Result<(), String> {
let (x, y) = resolve_element_center(client, session_id, ref_map, selector_or_ref).await?;
dispatch_click(client, session_id, x, y, button, click_count).await
}
pub async fn dblclick(
client: &CdpClient,
session_id: &str,
ref_map: &RefMap,
selector_or_ref: &str,
) -> Result<(), String> {
click(client, session_id, ref_map, selector_or_ref, "left", 2).await
}
pub async fn hover(
client: &CdpClient,
session_id: &str,
ref_map: &RefMap,
selector_or_ref: &str,
) -> Result<(), String> {
let (x, y) = resolve_element_center(client, session_id, ref_map, selector_or_ref).await?;
client
.send_command_typed::<_, Value>(
"Input.dispatchMouseEvent",
&DispatchMouseEventParams {
event_type: "mouseMoved".to_string(),
x,
y,
button: None,
buttons: None,
click_count: None,
delta_x: None,
delta_y: None,
modifiers: None,
},
Some(session_id),
)
.await?;
Ok(())
}
pub async fn fill(
client: &CdpClient,
session_id: &str,
ref_map: &RefMap,
selector_or_ref: &str,
value: &str,
) -> Result<(), String> {
let object_id = resolve_element_object_id(client, session_id, ref_map, selector_or_ref).await?;
// Focus the element
client
.send_command_typed::<_, Value>(
"Runtime.callFunctionOn",
&CallFunctionOnParams {
function_declaration: "function() { this.focus(); }".to_string(),
object_id: Some(object_id.clone()),
arguments: None,
return_by_value: Some(true),
await_promise: Some(false),
},
Some(session_id),
)
.await?;
// Select all + delete to clear
client
.send_command_typed::<_, Value>(
"Runtime.callFunctionOn",
&CallFunctionOnParams {
function_declaration: r#"function() {
this.select && this.select();
this.value = '';
this.dispatchEvent(new Event('input', { bubbles: true }));
}"#
.to_string(),
object_id: Some(object_id),
arguments: None,
return_by_value: Some(true),
await_promise: Some(false),
},
Some(session_id),
)
.await?;
// Insert text
client
.send_command_typed::<_, Value>(
"Input.insertText",
&InsertTextParams {
text: value.to_string(),
},
Some(session_id),
)
.await?;
Ok(())
}
pub async fn type_text(
client: &CdpClient,
session_id: &str,
ref_map: &RefMap,
selector_or_ref: &str,
text: &str,
clear: bool,
delay_ms: Option<u64>,
) -> Result<(), String> {
let object_id = resolve_element_object_id(client, session_id, ref_map, selector_or_ref).await?;
// Focus
client
.send_command_typed::<_, Value>(
"Runtime.callFunctionOn",
&CallFunctionOnParams {
function_declaration: "function() { this.focus(); }".to_string(),
object_id: Some(object_id.clone()),
arguments: None,
return_by_value: Some(true),
await_promise: Some(false),
},
Some(session_id),
)
.await?;
if clear {
client
.send_command_typed::<_, Value>(
"Runtime.callFunctionOn",
&CallFunctionOnParams {
function_declaration: r#"function() {
this.select && this.select();
this.value = '';
this.dispatchEvent(new Event('input', { bubbles: true }));
}"#
.to_string(),
object_id: Some(object_id),
arguments: None,
return_by_value: Some(true),
await_promise: Some(false),
},
Some(session_id),
)
.await?;
}
let delay = delay_ms.unwrap_or(0);
for ch in text.chars() {
let text_str = ch.to_string();
let (key, code, key_code) = char_to_key_info(ch);
client
.send_command_typed::<_, Value>(
"Input.dispatchKeyEvent",
&DispatchKeyEventParams {
event_type: "keyDown".to_string(),
key: Some(key.clone()),
code: Some(code.clone()),
text: Some(text_str.clone()),
unmodified_text: Some(text_str.clone()),
windows_virtual_key_code: Some(key_code),
native_virtual_key_code: Some(key_code),
modifiers: None,
},
Some(session_id),
)
.await?;
client
.send_command_typed::<_, Value>(
"Input.dispatchKeyEvent",
&DispatchKeyEventParams {
event_type: "keyUp".to_string(),
key: Some(key),
code: Some(code),
text: None,
unmodified_text: None,
windows_virtual_key_code: Some(key_code),
native_virtual_key_code: Some(key_code),
modifiers: None,
},
Some(session_id),
)
.await?;
if delay > 0 {
tokio::time::sleep(tokio::time::Duration::from_millis(delay)).await;
}
}
Ok(())
}
pub async fn press_key(client: &CdpClient, session_id: &str, key: &str) -> Result<(), String> {
let (key_name, code, key_code) = named_key_info(key);
client
.send_command_typed::<_, Value>(
"Input.dispatchKeyEvent",
&DispatchKeyEventParams {
event_type: "keyDown".to_string(),
key: Some(key_name.clone()),
code: Some(code.clone()),
text: None,
unmodified_text: None,
windows_virtual_key_code: Some(key_code),
native_virtual_key_code: Some(key_code),
modifiers: None,
},
Some(session_id),
)
.await?;
client
.send_command_typed::<_, Value>(
"Input.dispatchKeyEvent",
&DispatchKeyEventParams {
event_type: "keyUp".to_string(),
key: Some(key_name),
code: Some(code),
text: None,
unmodified_text: None,
windows_virtual_key_code: Some(key_code),
native_virtual_key_code: Some(key_code),
modifiers: None,
},
Some(session_id),
)
.await?;
Ok(())
}
pub async fn scroll(
client: &CdpClient,
session_id: &str,
ref_map: &RefMap,
selector_or_ref: Option<&str>,
delta_x: f64,
delta_y: f64,
) -> Result<(), String> {
if let Some(sel) = selector_or_ref {
let object_id = resolve_element_object_id(client, session_id, ref_map, sel).await?;
let js = "function(dx, dy) { this.scrollBy(dx, dy); }".to_string();
client
.send_command_typed::<_, Value>(
"Runtime.callFunctionOn",
&CallFunctionOnParams {
function_declaration: js,
object_id: Some(object_id),
arguments: Some(vec![
CallArgument {
value: Some(serde_json::json!(delta_x)),
object_id: None,
},
CallArgument {
value: Some(serde_json::json!(delta_y)),
object_id: None,
},
]),
return_by_value: Some(true),
await_promise: Some(false),
},
Some(session_id),
)
.await?;
} else {
let js = format!("window.scrollBy({}, {})", delta_x, delta_y);
client
.send_command_typed::<_, Value>(
"Runtime.evaluate",
&EvaluateParams {
expression: js,
return_by_value: Some(true),
await_promise: Some(false),
},
Some(session_id),
)
.await?;
}
Ok(())
}
pub async fn select_option(
client: &CdpClient,
session_id: &str,
ref_map: &RefMap,
selector_or_ref: &str,
values: &[String],
) -> Result<(), String> {
let object_id = resolve_element_object_id(client, session_id, ref_map, selector_or_ref).await?;
let js = r#"function(vals) {
const options = Array.from(this.options);
for (const opt of options) {
opt.selected = vals.includes(opt.value) || vals.includes(opt.textContent.trim());
}
this.dispatchEvent(new Event('change', { bubbles: true }));
}"#
.to_string();
client
.send_command_typed::<_, Value>(
"Runtime.callFunctionOn",
&CallFunctionOnParams {
function_declaration: js,
object_id: Some(object_id),
arguments: Some(vec![CallArgument {
value: Some(serde_json::json!(values)),
object_id: None,
}]),
return_by_value: Some(true),
await_promise: Some(false),
},
Some(session_id),
)
.await?;
Ok(())
}
pub async fn check(
client: &CdpClient,
session_id: &str,
ref_map: &RefMap,
selector_or_ref: &str,
) -> Result<(), String> {
let is_checked =
super::element::is_element_checked(client, session_id, ref_map, selector_or_ref).await?;
if !is_checked {
click(client, session_id, ref_map, selector_or_ref, "left", 1).await?;
}
Ok(())
}
pub async fn uncheck(
client: &CdpClient,
session_id: &str,
ref_map: &RefMap,
selector_or_ref: &str,
) -> Result<(), String> {
let is_checked =
super::element::is_element_checked(client, session_id, ref_map, selector_or_ref).await?;
if is_checked {
click(client, session_id, ref_map, selector_or_ref, "left", 1).await?;
}
Ok(())
}
pub async fn focus(
client: &CdpClient,
session_id: &str,
ref_map: &RefMap,
selector_or_ref: &str,
) -> Result<(), String> {
let object_id = resolve_element_object_id(client, session_id, ref_map, selector_or_ref).await?;
client
.send_command_typed::<_, Value>(
"Runtime.callFunctionOn",
&CallFunctionOnParams {
function_declaration: "function() { this.focus(); }".to_string(),
object_id: Some(object_id),
arguments: None,
return_by_value: Some(true),
await_promise: Some(false),
},
Some(session_id),
)
.await?;
Ok(())
}
pub async fn clear(
client: &CdpClient,
session_id: &str,
ref_map: &RefMap,
selector_or_ref: &str,
) -> Result<(), String> {
let object_id = resolve_element_object_id(client, session_id, ref_map, selector_or_ref).await?;
client
.send_command_typed::<_, Value>(
"Runtime.callFunctionOn",
&CallFunctionOnParams {
function_declaration: r#"function() {
this.focus();
this.value = '';
this.dispatchEvent(new Event('input', { bubbles: true }));
this.dispatchEvent(new Event('change', { bubbles: true }));
}"#
.to_string(),
object_id: Some(object_id),
arguments: None,
return_by_value: Some(true),
await_promise: Some(false),
},
Some(session_id),
)
.await?;
Ok(())
}
pub async fn select_all(
client: &CdpClient,
session_id: &str,
ref_map: &RefMap,
selector_or_ref: &str,
) -> Result<(), String> {
let object_id = resolve_element_object_id(client, session_id, ref_map, selector_or_ref).await?;
client
.send_command_typed::<_, Value>(
"Runtime.callFunctionOn",
&CallFunctionOnParams {
function_declaration: r#"function() {
this.focus();
if (typeof this.select === 'function') {
this.select();
} else {
const range = document.createRange();
range.selectNodeContents(this);
const sel = window.getSelection();
sel.removeAllRanges();
sel.addRange(range);
}
}"#
.to_string(),
object_id: Some(object_id),
arguments: None,
return_by_value: Some(true),
await_promise: Some(false),
},
Some(session_id),
)
.await?;
Ok(())
}
pub async fn scroll_into_view(
client: &CdpClient,
session_id: &str,
ref_map: &RefMap,
selector_or_ref: &str,
) -> Result<(), String> {
let object_id = resolve_element_object_id(client, session_id, ref_map, selector_or_ref).await?;
client
.send_command_typed::<_, Value>(
"Runtime.callFunctionOn",
&CallFunctionOnParams {
function_declaration:
"function() { this.scrollIntoView({ block: 'center', inline: 'center' }); }"
.to_string(),
object_id: Some(object_id),
arguments: None,
return_by_value: Some(true),
await_promise: Some(false),
},
Some(session_id),
)
.await?;
Ok(())
}
pub async fn dispatch_event(
client: &CdpClient,
session_id: &str,
ref_map: &RefMap,
selector_or_ref: &str,
event_type: &str,
event_init: Option<&Value>,
) -> Result<(), String> {
let object_id = resolve_element_object_id(client, session_id, ref_map, selector_or_ref).await?;
let init_json = event_init
.map(|v| serde_json::to_string(v).unwrap_or("{}".to_string()))
.unwrap_or_else(|| "{ bubbles: true }".to_string());
let js = format!(
"function() {{ this.dispatchEvent(new Event({}, {})); }}",
serde_json::to_string(event_type).unwrap_or_default(),
init_json
);
client
.send_command_typed::<_, Value>(
"Runtime.callFunctionOn",
&CallFunctionOnParams {
function_declaration: js,
object_id: Some(object_id),
arguments: None,
return_by_value: Some(true),
await_promise: Some(false),
},
Some(session_id),
)
.await?;
Ok(())
}
pub async fn highlight(
client: &CdpClient,
session_id: &str,
ref_map: &RefMap,
selector_or_ref: &str,
) -> Result<(), String> {
let object_id = resolve_element_object_id(client, session_id, ref_map, selector_or_ref).await?;
client
.send_command_typed::<_, Value>(
"Runtime.callFunctionOn",
&CallFunctionOnParams {
function_declaration: r#"function() {
this.style.outline = '2px solid red';
this.style.outlineOffset = '2px';
const el = this;
setTimeout(() => {
el.style.outline = '';
el.style.outlineOffset = '';
}, 3000);
}"#
.to_string(),
object_id: Some(object_id),
arguments: None,
return_by_value: Some(true),
await_promise: Some(false),
},
Some(session_id),
)
.await?;
Ok(())
}
pub async fn tap_touch(
client: &CdpClient,
session_id: &str,
ref_map: &RefMap,
selector_or_ref: &str,
) -> Result<(), String> {
let (x, y) = resolve_element_center(client, session_id, ref_map, selector_or_ref).await?;
client
.send_command(
"Input.dispatchTouchEvent",
Some(serde_json::json!({
"type": "touchStart",
"touchPoints": [{ "x": x, "y": y }],
})),
Some(session_id),
)
.await?;
client
.send_command(
"Input.dispatchTouchEvent",
Some(serde_json::json!({
"type": "touchEnd",
"touchPoints": [],
})),
Some(session_id),
)
.await?;
Ok(())
}
async fn dispatch_click(
client: &CdpClient,
session_id: &str,
x: f64,
y: f64,
button: &str,
click_count: i32,
) -> Result<(), String> {
// Move
client
.send_command_typed::<_, Value>(
"Input.dispatchMouseEvent",
&DispatchMouseEventParams {
event_type: "mouseMoved".to_string(),
x,
y,
button: None,
buttons: None,
click_count: None,
delta_x: None,
delta_y: None,
modifiers: None,
},
Some(session_id),
)
.await?;
let button_value = match button {
"right" => 2,
"middle" => 4,
_ => 1,
};
// Press
client
.send_command_typed::<_, Value>(
"Input.dispatchMouseEvent",
&DispatchMouseEventParams {
event_type: "mousePressed".to_string(),
x,
y,
button: Some(button.to_string()),
buttons: Some(button_value),
click_count: Some(click_count),
delta_x: None,
delta_y: None,
modifiers: None,
},
Some(session_id),
)
.await?;
// Release
client
.send_command_typed::<_, Value>(
"Input.dispatchMouseEvent",
&DispatchMouseEventParams {
event_type: "mouseReleased".to_string(),
x,
y,
button: Some(button.to_string()),
buttons: Some(0),
click_count: Some(click_count),
delta_x: None,
delta_y: None,
modifiers: None,
},
Some(session_id),
)
.await?;
Ok(())
}
fn char_to_key_info(ch: char) -> (String, String, i32) {
match ch {
'\n' | '\r' => ("Enter".to_string(), "Enter".to_string(), 13),
'\t' => ("Tab".to_string(), "Tab".to_string(), 9),
' ' => (" ".to_string(), "Space".to_string(), 32),
_ => {
let key = ch.to_string();
let code = if ch.is_ascii_alphabetic() {
format!("Key{}", ch.to_uppercase())
} else if ch.is_ascii_digit() {
format!("Digit{}", ch)
} else {
String::new()
};
let key_code = ch as i32;
(key, code, key_code)
}
}
}
fn named_key_info(key: &str) -> (String, String, i32) {
match key.to_lowercase().as_str() {
"enter" | "return" => ("Enter".to_string(), "Enter".to_string(), 13),
"tab" => ("Tab".to_string(), "Tab".to_string(), 9),
"escape" | "esc" => ("Escape".to_string(), "Escape".to_string(), 27),
"backspace" => ("Backspace".to_string(), "Backspace".to_string(), 8),
"delete" => ("Delete".to_string(), "Delete".to_string(), 46),
"arrowup" | "up" => ("ArrowUp".to_string(), "ArrowUp".to_string(), 38),
"arrowdown" | "down" => ("ArrowDown".to_string(), "ArrowDown".to_string(), 40),
"arrowleft" | "left" => ("ArrowLeft".to_string(), "ArrowLeft".to_string(), 37),
"arrowright" | "right" => ("ArrowRight".to_string(), "ArrowRight".to_string(), 39),
"home" => ("Home".to_string(), "Home".to_string(), 36),
"end" => ("End".to_string(), "End".to_string(), 35),
"pageup" => ("PageUp".to_string(), "PageUp".to_string(), 33),
"pagedown" => ("PageDown".to_string(), "PageDown".to_string(), 34),
"space" | " " => (" ".to_string(), "Space".to_string(), 32),
_ => {
if key.len() == 1 {
let ch = key.chars().next().unwrap();
char_to_key_info(ch)
} else {
(key.to_string(), key.to_string(), 0)
}
}
}
}
+45
View File
@@ -0,0 +1,45 @@
#[allow(dead_code)]
pub mod actions;
#[allow(dead_code)]
pub mod auth;
#[allow(dead_code)]
pub mod browser;
#[allow(dead_code)]
pub mod cdp;
#[allow(dead_code)]
pub mod cookies;
#[allow(dead_code)]
pub mod daemon;
#[allow(dead_code)]
pub mod diff;
#[allow(dead_code)]
pub mod element;
#[allow(dead_code)]
pub mod interaction;
#[allow(dead_code)]
pub mod network;
#[allow(dead_code)]
pub mod policy;
#[allow(dead_code)]
pub mod providers;
#[allow(dead_code)]
pub mod recording;
#[allow(dead_code)]
pub mod screenshot;
#[allow(dead_code)]
pub mod snapshot;
#[allow(dead_code)]
pub mod state;
#[allow(dead_code)]
pub mod storage;
#[allow(dead_code)]
pub mod stream;
#[allow(dead_code)]
pub mod tracing;
#[allow(dead_code)]
pub mod webdriver;
#[cfg(test)]
mod e2e_tests;
#[cfg(test)]
mod parity_tests;
+399
View File
@@ -0,0 +1,399 @@
use serde_json::{json, Value};
use std::collections::HashMap;
use super::cdp::client::CdpClient;
pub async fn set_extra_headers(
client: &CdpClient,
session_id: &str,
headers: &HashMap<String, String>,
) -> Result<(), String> {
let headers_value: Value = headers
.iter()
.map(|(k, v)| (k.clone(), Value::String(v.clone())))
.collect::<serde_json::Map<String, Value>>()
.into();
client
.send_command(
"Network.setExtraHTTPHeaders",
Some(json!({ "headers": headers_value })),
Some(session_id),
)
.await?;
Ok(())
}
pub async fn set_offline(
client: &CdpClient,
session_id: &str,
offline: bool,
) -> Result<(), String> {
client
.send_command(
"Network.emulateNetworkConditions",
Some(json!({
"offline": offline,
"latency": 0,
"downloadThroughput": -1,
"uploadThroughput": -1,
})),
Some(session_id),
)
.await?;
Ok(())
}
pub async fn set_content(client: &CdpClient, session_id: &str, html: &str) -> Result<(), String> {
// Get current frame ID
let tree_result = client
.send_command_no_params("Page.getFrameTree", Some(session_id))
.await?;
let frame_id = tree_result
.get("frameTree")
.and_then(|t| t.get("frame"))
.and_then(|f| f.get("id"))
.and_then(|id| id.as_str())
.ok_or("Could not determine frame ID")?;
client
.send_command(
"Page.setDocumentContent",
Some(json!({
"frameId": frame_id,
"html": html,
})),
Some(session_id),
)
.await?;
Ok(())
}
// ---------------------------------------------------------------------------
// Domain filter
// ---------------------------------------------------------------------------
#[derive(Debug, Clone)]
pub struct DomainFilter {
pub allowed_domains: Vec<String>,
}
impl DomainFilter {
pub fn new(domains: &str) -> Self {
let allowed = parse_domain_list(domains);
Self {
allowed_domains: allowed,
}
}
pub fn is_allowed(&self, hostname: &str) -> bool {
if self.allowed_domains.is_empty() {
return true;
}
let hostname = hostname.to_lowercase();
for pattern in &self.allowed_domains {
if let Some(suffix) = pattern.strip_prefix("*.") {
if hostname == suffix || hostname.ends_with(&format!(".{}", suffix)) {
return true;
}
} else if hostname == *pattern {
return true;
}
}
false
}
pub fn check_url(&self, url: &str) -> Result<(), String> {
if self.allowed_domains.is_empty() {
return Ok(());
}
let parsed = url::Url::parse(url).map_err(|_| format!("Invalid URL: {}", url))?;
let hostname = parsed
.host_str()
.ok_or_else(|| format!("No hostname in URL: {}", url))?;
if self.is_allowed(hostname) {
Ok(())
} else {
Err(format!(
"Domain '{}' is not in the allowed domains list",
hostname
))
}
}
}
fn parse_domain_list(input: &str) -> Vec<String> {
input
.split(',')
.map(|s| s.trim().to_lowercase())
.filter(|s| !s.is_empty())
.collect()
}
pub async fn sanitize_existing_pages(
client: &CdpClient,
pages: &[super::browser::PageInfo],
filter: &DomainFilter,
) {
for page in pages {
if page.url.is_empty() || page.url == "about:blank" {
continue;
}
if let Ok(parsed) = url::Url::parse(&page.url) {
if let Some(hostname) = parsed.host_str() {
if !filter.is_allowed(hostname) {
let _ = client
.send_command(
"Page.navigate",
Some(json!({ "url": "about:blank" })),
Some(&page.session_id),
)
.await;
}
}
}
}
}
pub async fn install_domain_filter_script(
client: &CdpClient,
session_id: &str,
allowed_domains: &[String],
) -> Result<(), String> {
if allowed_domains.is_empty() {
return Ok(());
}
let domains_json = serde_json::to_string(allowed_domains).unwrap_or("[]".to_string());
let script = format!(
r#"(() => {{
const _allowed = {};
function _isDomainAllowed(hostname) {{
hostname = hostname.toLowerCase();
for (const p of _allowed) {{
if (p.startsWith('*.')) {{
const suffix = p.slice(2);
if (hostname === suffix || hostname.endsWith('.' + suffix)) return true;
}} else if (hostname === p) return true;
}}
return false;
}}
const OrigWS = window.WebSocket;
window.WebSocket = function(url, protocols) {{
try {{
const u = new URL(url);
if (!_isDomainAllowed(u.hostname)) throw new DOMException('WebSocket blocked: ' + u.hostname, 'SecurityError');
}} catch(e) {{ if (e instanceof DOMException) throw e; }}
return new OrigWS(url, protocols);
}};
window.WebSocket.prototype = OrigWS.prototype;
const OrigES = window.EventSource;
if (OrigES) {{
window.EventSource = function(url, opts) {{
try {{
const u = new URL(url, location.href);
if (!_isDomainAllowed(u.hostname)) throw new DOMException('EventSource blocked: ' + u.hostname, 'SecurityError');
}} catch(e) {{ if (e instanceof DOMException) throw e; }}
return new OrigES(url, opts);
}};
window.EventSource.prototype = OrigES.prototype;
}}
const origBeacon = navigator.sendBeacon;
if (origBeacon) {{
navigator.sendBeacon = function(url, data) {{
try {{
const u = new URL(url, location.href);
if (!_isDomainAllowed(u.hostname)) return false;
}} catch(e) {{ return false; }}
return origBeacon.call(navigator, url, data);
}};
}}
}})()"#,
domains_json,
);
client
.send_command(
"Page.addScriptToEvaluateOnNewDocument",
Some(json!({ "source": script })),
Some(session_id),
)
.await?;
Ok(())
}
/// Enable Fetch-based network interception for domain filtering.
/// This intercepts all requests and checks them against the allowed domains list.
/// The actual handling of `Fetch.requestPaused` events happens in
/// `resolve_fetch_paused` in the actions module.
pub async fn install_domain_filter_fetch(
client: &CdpClient,
session_id: &str,
) -> Result<(), String> {
client
.send_command(
"Fetch.enable",
Some(json!({
"patterns": [{ "urlPattern": "*" }]
})),
Some(session_id),
)
.await?;
Ok(())
}
/// Install both layers of domain filtering on a session:
/// 1. JS patching (WebSocket, EventSource, sendBeacon)
/// 2. Fetch-based network interception
pub async fn install_domain_filter(
client: &CdpClient,
session_id: &str,
allowed_domains: &[String],
) -> Result<(), String> {
install_domain_filter_script(client, session_id, allowed_domains).await?;
install_domain_filter_fetch(client, session_id).await?;
Ok(())
}
// ---------------------------------------------------------------------------
// Console and error tracking
// ---------------------------------------------------------------------------
#[derive(Debug, Clone)]
pub struct ConsoleEntry {
pub level: String,
pub text: String,
}
#[derive(Debug, Clone)]
pub struct ErrorEntry {
pub text: String,
pub url: Option<String>,
pub line: Option<i64>,
pub column: Option<i64>,
}
pub struct EventTracker {
pub console_entries: Vec<ConsoleEntry>,
pub error_entries: Vec<ErrorEntry>,
pub max_entries: usize,
}
impl EventTracker {
pub fn new() -> Self {
Self {
console_entries: Vec::new(),
error_entries: Vec::new(),
max_entries: 1000,
}
}
pub fn add_console(&mut self, level: &str, text: &str) {
if self.console_entries.len() >= self.max_entries {
self.console_entries.remove(0);
}
self.console_entries.push(ConsoleEntry {
level: level.to_string(),
text: text.to_string(),
});
}
pub fn add_error(
&mut self,
text: &str,
url: Option<&str>,
line: Option<i64>,
col: Option<i64>,
) {
if self.error_entries.len() >= self.max_entries {
self.error_entries.remove(0);
}
self.error_entries.push(ErrorEntry {
text: text.to_string(),
url: url.map(String::from),
line,
column: col,
});
}
pub fn get_console_json(&self) -> Value {
let entries: Vec<Value> = self
.console_entries
.iter()
.map(|e| json!({ "level": e.level, "text": e.text }))
.collect();
json!({ "entries": entries })
}
pub fn get_errors_json(&self) -> Value {
let entries: Vec<Value> = self
.error_entries
.iter()
.map(|e| {
json!({
"text": e.text,
"url": e.url,
"line": e.line,
"column": e.column,
})
})
.collect();
json!({ "errors": entries })
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_domain_filter_exact() {
let filter = DomainFilter::new("example.com");
assert!(filter.is_allowed("example.com"));
assert!(!filter.is_allowed("other.com"));
}
#[test]
fn test_domain_filter_wildcard() {
let filter = DomainFilter::new("*.example.com");
assert!(filter.is_allowed("example.com"));
assert!(filter.is_allowed("api.example.com"));
assert!(filter.is_allowed("sub.api.example.com"));
assert!(!filter.is_allowed("other.com"));
}
#[test]
fn test_domain_filter_empty() {
let filter = DomainFilter::new("");
assert!(filter.is_allowed("anything.com"));
}
#[test]
fn test_domain_filter_multiple() {
let filter = DomainFilter::new("example.com, *.api.io");
assert!(filter.is_allowed("example.com"));
assert!(filter.is_allowed("api.io"));
assert!(filter.is_allowed("v1.api.io"));
assert!(!filter.is_allowed("other.com"));
}
#[test]
fn test_parse_domain_list() {
let domains = parse_domain_list("A.com, B.com , *.C.com");
assert_eq!(domains, vec!["a.com", "b.com", "*.c.com"]);
}
#[test]
fn test_event_tracker() {
let mut tracker = EventTracker::new();
tracker.add_console("log", "hello");
tracker.add_error("oops", Some("test.js"), Some(1), Some(5));
assert_eq!(tracker.console_entries.len(), 1);
assert_eq!(tracker.error_entries.len(), 1);
}
}
+625
View File
@@ -0,0 +1,625 @@
//! Parity tests for the native daemon's command interface.
//!
//! These unit tests verify:
//! - All documented actions are handled (not returning "Not yet implemented")
//! - Response format consistency (success/error structure)
//! - Credential and state actions work without a browser
use serde_json::{json, Value};
use super::actions::{execute_command, DaemonState};
/// All documented action names that should be implemented.
const DOCUMENTED_ACTIONS: &[&str] = &[
"launch",
"navigate",
"url",
"title",
"content",
"evaluate",
"close",
"snapshot",
"screenshot",
"click",
"dblclick",
"fill",
"type",
"press",
"hover",
"scroll",
"select",
"check",
"uncheck",
"wait",
"gettext",
"getattribute",
"isvisible",
"isenabled",
"ischecked",
"back",
"forward",
"reload",
"cookies_get",
"cookies_set",
"cookies_clear",
"storage_get",
"storage_set",
"storage_clear",
"setcontent",
"headers",
"offline",
"console",
"errors",
"state_save",
"state_load",
"state_list",
"state_show",
"state_clear",
"state_clean",
"state_rename",
"trace_start",
"trace_stop",
"profiler_start",
"profiler_stop",
"recording_start",
"recording_stop",
"recording_restart",
"pdf",
"tab_list",
"tab_new",
"tab_switch",
"tab_close",
"viewport",
"user_agent",
"set_media",
"download",
"diff_snapshot",
"diff_url",
"credentials_set",
"credentials_get",
"credentials_delete",
"credentials_list",
"mouse",
"keyboard",
"focus",
"clear",
"selectall",
"scrollintoview",
"dispatch",
"highlight",
"tap",
"boundingbox",
"innertext",
"innerhtml",
"inputvalue",
"setvalue",
"count",
"styles",
"bringtofront",
"timezone",
"locale",
"geolocation",
"permissions",
"dialog",
"upload",
"addscript",
"addinitscript",
"addstyle",
"clipboard",
"wheel",
"device",
"screencast_start",
"screencast_stop",
"waitforurl",
"waitforloadstate",
"waitforfunction",
"frame",
"mainframe",
"getbyrole",
"getbytext",
"getbylabel",
"getbyplaceholder",
"getbyalttext",
"getbytitle",
"getbytestid",
"nth",
"find",
"evalhandle",
"drag",
"expose",
"pause",
"multiselect",
"responsebody",
"waitfordownload",
"window_new",
"diff_screenshot",
"video_start",
"video_stop",
"har_start",
"har_stop",
"route",
"unroute",
"requests",
"credentials",
"auth_save",
"auth_login",
"auth_list",
"auth_delete",
"auth_show",
"confirm",
"deny",
"swipe",
"device_list",
"input_mouse",
"input_keyboard",
"input_touch",
"keydown",
"keyup",
"inserttext",
"mousemove",
"mousedown",
"mouseup",
];
fn minimal_command(action: &str, id: &str) -> Value {
let mut cmd = json!({ "action": action, "id": id });
let obj = cmd.as_object_mut().unwrap();
match action {
"navigate" | "diff_url" | "waitforurl" => {
obj.insert("url".to_string(), json!("https://example.com"));
}
"evaluate" | "expose" => {
obj.insert("script".to_string(), json!("1"));
}
"click" | "dblclick" | "fill" | "type" | "press" | "hover" | "scroll" | "select"
| "check" | "uncheck" | "gettext" | "getattribute" | "isvisible" | "isenabled"
| "ischecked" | "focus" | "clear" | "selectall" | "scrollintoview" | "dispatch"
| "highlight" | "tap" | "boundingbox" | "innertext" | "innerhtml" | "inputvalue"
| "setvalue" | "count" | "find" | "nth" | "getbytext" | "getbylabel"
| "getbyplaceholder" | "getbyalttext" | "getbytitle" | "getbytestid" => {
obj.insert("selector".to_string(), json!("body"));
}
"getbyrole" => {
obj.insert("role".to_string(), json!("button"));
obj.insert("selector".to_string(), json!("body"));
}
"setcontent" => {
obj.insert("html".to_string(), json!("<html></html>"));
}
"cookies_set" => {
obj.insert("name".to_string(), json!("test"));
obj.insert("value".to_string(), json!("val"));
}
"storage_get" | "storage_set" | "storage_clear" => {
obj.insert("origin".to_string(), json!("https://example.com"));
}
"state_save" | "state_load" | "state_show" | "state_clear" => {
obj.insert("path".to_string(), json!("test-parity-state.json"));
}
"state_rename" => {
obj.insert("path".to_string(), json!("test-parity-state.json"));
obj.insert("name".to_string(), json!("renamed"));
}
"state_clean" => {
obj.insert("days".to_string(), json!(7));
}
"credentials_set" => {
obj.insert("name".to_string(), json!("parity-test-cred"));
obj.insert("username".to_string(), json!("u"));
obj.insert("password".to_string(), json!("p"));
}
"auth_save" => {
obj.insert("name".to_string(), json!("parity-test-cred"));
obj.insert("url".to_string(), json!("https://example.com"));
obj.insert("username".to_string(), json!("u"));
obj.insert("password".to_string(), json!("p"));
}
"credentials_get" | "credentials_delete" | "auth_show" | "auth_delete" => {
obj.insert("name".to_string(), json!("parity-test-cred"));
}
"tab_switch" | "tab_close" => {
obj.insert("index".to_string(), json!(0));
}
"viewport" | "user_agent" | "set_media" | "timezone" | "locale" | "geolocation"
| "permissions" | "device" => {
obj.insert("value".to_string(), json!(null));
}
"headers" => {
obj.insert("headers".to_string(), json!({}));
}
"offline" => {
obj.insert("offline".to_string(), json!(false));
}
"wait" => {
obj.insert("timeout".to_string(), json!(100));
}
"waitforloadstate" => {
obj.insert("state".to_string(), json!("load"));
}
"waitforfunction" => {
obj.insert("script".to_string(), json!("() => true"));
}
"frame" => {
obj.insert("selector".to_string(), json!("iframe"));
}
"addscript" => {
obj.insert("content".to_string(), json!("console.log('test')"));
}
"addinitscript" => {
obj.insert("script".to_string(), json!("console.log('init')"));
}
"addstyle" => {
obj.insert("content".to_string(), json!("body { color: red }"));
}
"wheel" => {
obj.insert("deltaX".to_string(), json!(0));
obj.insert("deltaY".to_string(), json!(0));
}
"upload" => {
obj.insert("selector".to_string(), json!("input[type=file]"));
obj.insert("files".to_string(), json!([]));
}
"dialog" => {
obj.insert("accept".to_string(), json!(true));
}
"credentials" => {
obj.insert("username".to_string(), json!("u"));
obj.insert("password".to_string(), json!("p"));
}
"auth_login" => {
obj.insert("name".to_string(), json!("parity-test-cred"));
}
"route" => {
obj.insert("url".to_string(), json!("*"));
obj.insert("handler".to_string(), json!("continue"));
}
"diff_snapshot" | "diff_screenshot" => {
obj.insert("selector".to_string(), json!("body"));
}
"recording_start" | "recording_restart" => {
obj.insert("path".to_string(), json!("/tmp/parity-recording.webm"));
}
"video_start" => {
obj.insert("path".to_string(), json!("/tmp/parity-video.webm"));
}
"profiler_start" => {
obj.insert("path".to_string(), json!("/tmp/parity-profile"));
}
"trace_stop" | "har_stop" => {
obj.insert("path".to_string(), json!("/tmp/parity-trace"));
}
"download" => {
obj.insert("path".to_string(), json!("/tmp/parity-download"));
}
"multiselect" => {
obj.insert("selector".to_string(), json!("select"));
obj.insert("values".to_string(), json!([]));
}
"responsebody" => {
obj.insert("url".to_string(), json!("https://example.com"));
}
"waitfordownload" => {
obj.insert("path".to_string(), json!("/tmp/parity-download"));
}
"styles" => {
obj.insert("selector".to_string(), json!("body"));
obj.insert("names".to_string(), json!([]));
}
"evalhandle" => {
obj.insert("handle".to_string(), json!(""));
obj.insert("script".to_string(), json!("h => h"));
}
"drag" => {
obj.insert("selector".to_string(), json!("body"));
obj.insert("target".to_string(), json!("body"));
}
"swipe" => {
obj.insert("selector".to_string(), json!("body"));
obj.insert("direction".to_string(), json!("left"));
}
"input_mouse" | "mousemove" | "mousedown" | "mouseup" => {
obj.insert("x".to_string(), json!(100));
obj.insert("y".to_string(), json!(100));
}
"input_keyboard" | "keydown" | "keyup" => {
obj.insert("key".to_string(), json!("a"));
}
"input_touch" => {
obj.insert("type".to_string(), json!("touchStart"));
obj.insert("touchPoints".to_string(), json!([]));
}
"inserttext" => {
obj.insert("text".to_string(), json!("test"));
}
_ => {}
}
cmd
}
// ---------------------------------------------------------------------------
// 1. Action dispatch coverage
// ---------------------------------------------------------------------------
#[tokio::test]
async fn test_all_documented_actions_are_handled() {
let mut state = DaemonState::new();
for (i, action) in DOCUMENTED_ACTIONS.iter().enumerate() {
let id = format!("parity-{}", i);
let cmd = minimal_command(action, &id);
let result = execute_command(&cmd, &mut state).await;
assert!(
result.get("id").is_some(),
"Action '{}': response missing 'id'",
action
);
let error = result.get("error").and_then(|v| v.as_str()).unwrap_or("");
assert!(
!error.contains("Not yet implemented"),
"Action '{}' returned 'Not yet implemented')",
action
);
}
}
// ---------------------------------------------------------------------------
// 2. Response format consistency
// ---------------------------------------------------------------------------
#[tokio::test]
async fn test_success_response_format() {
let mut state = DaemonState::new();
let cmd = json!({ "action": "state_list", "id": "fmt-1" });
let result = execute_command(&cmd, &mut state).await;
assert_eq!(result["success"], true);
assert!(result.get("id").is_some());
assert!(result.get("data").is_some());
assert!(result.get("error").is_none());
}
#[tokio::test]
async fn test_error_response_format() {
let mut state = DaemonState::new();
let cmd = json!({ "action": "nonexistent_action_xyz", "id": "fmt-2" });
let result = execute_command(&cmd, &mut state).await;
assert_eq!(result["success"], false);
assert!(result.get("id").is_some());
assert!(result.get("error").is_some());
}
// ---------------------------------------------------------------------------
// 3. Credential/state actions work without a browser
// ---------------------------------------------------------------------------
#[tokio::test]
async fn test_state_list_without_browser() {
let mut state = DaemonState::new();
let cmd = json!({ "action": "state_list", "id": "nb-1" });
let result = execute_command(&cmd, &mut state).await;
assert_eq!(result["success"], true);
assert!(result["data"]["files"].is_array());
}
#[tokio::test]
async fn test_credentials_list_without_browser() {
let mut state = DaemonState::new();
let cmd = json!({ "action": "credentials_list", "id": "nb-2" });
let result = execute_command(&cmd, &mut state).await;
assert_eq!(result["success"], true);
assert!(result["data"]["credentials"].is_array() || result["data"]["profiles"].is_array());
}
// ---------------------------------------------------------------------------
// 4. New feature parity tests
// ---------------------------------------------------------------------------
#[tokio::test]
async fn test_auth_profile_name_validation() {
use super::auth;
let valid = auth::credentials_set("valid-name_123", "u", "p", None);
assert!(valid.is_ok());
let invalid = auth::credentials_set("invalid/name", "u", "p", None);
assert!(invalid.is_err());
let invalid2 = auth::credentials_set("", "u", "p", None);
assert!(invalid2.is_err());
let invalid3 = auth::credentials_set("has space", "u", "p", None);
assert!(invalid3.is_err());
// Cleanup
let _ = auth::credentials_delete("valid-name_123");
}
#[tokio::test]
async fn test_auth_save_and_show() {
use super::auth;
let result = auth::auth_save(
"parity-roundtrip",
"https://example.com",
"user",
"pass",
Some("input#user"),
None,
None,
);
assert!(result.is_ok());
let show = auth::auth_show("parity-roundtrip");
assert!(show.is_ok());
let data = show.unwrap();
assert_eq!(data["profile"]["username"], "user");
assert_eq!(data["profile"]["usernameSelector"], "input#user");
let full = auth::credentials_get_full("parity-roundtrip");
assert!(full.is_ok());
assert_eq!(full.unwrap().password, "pass");
// Cleanup
let _ = auth::credentials_delete("parity-roundtrip");
}
#[tokio::test]
async fn test_har_start_stop_without_browser() {
let mut state = DaemonState::new();
// har_start requires a browser. Because execute_command auto-launches when
// no browser is present, the result depends on Chrome availability: success
// if Chrome is found (CI), failure if not. Both outcomes are valid.
let cmd = json!({ "action": "har_start", "id": "har-1" });
let result = execute_command(&cmd, &mut state).await;
let success = result["success"].as_bool().unwrap_or(false);
if success {
assert!(state.har_recording);
} else {
assert!(result["error"].as_str().is_some());
}
}
#[tokio::test]
async fn test_state_clean_action() {
let mut state = DaemonState::new();
let cmd = json!({ "action": "state_clean", "id": "clean-1", "days": 30 });
let result = execute_command(&cmd, &mut state).await;
assert_eq!(result["success"], true);
}
#[tokio::test]
async fn test_daemon_state_new_defaults() {
let state = DaemonState::new();
assert!(state.browser.is_none());
assert!(!state.har_recording);
assert!(state.har_entries.is_empty());
assert!(state.pending_confirmation.is_none());
assert!(!state.request_tracking);
assert!(state.tracked_requests.is_empty());
assert!(state.active_frame_id.is_none());
assert!(state.webdriver_backend.is_none());
}
#[tokio::test]
async fn test_tracked_request_struct() {
use super::actions::TrackedRequest;
let tr = TrackedRequest {
url: "https://example.com/api".to_string(),
method: "GET".to_string(),
headers: json!({"Accept": "text/html"}),
timestamp: 12345,
resource_type: "Document".to_string(),
};
let serialized = serde_json::to_value(&tr).unwrap();
assert_eq!(serialized["url"], "https://example.com/api");
assert_eq!(serialized["method"], "GET");
assert_eq!(serialized["resourceType"], "Document");
assert_eq!(serialized["timestamp"], 12345);
}
#[tokio::test]
async fn test_request_tracking_state() {
let mut state = DaemonState::new();
assert!(!state.request_tracking);
assert!(state.tracked_requests.is_empty());
state.tracked_requests.push(super::actions::TrackedRequest {
url: "https://example.com".to_string(),
method: "GET".to_string(),
headers: json!({}),
timestamp: 1,
resource_type: "Document".to_string(),
});
state.tracked_requests.push(super::actions::TrackedRequest {
url: "https://other.com".to_string(),
method: "POST".to_string(),
headers: json!({}),
timestamp: 2,
resource_type: "XHR".to_string(),
});
assert_eq!(state.tracked_requests.len(), 2);
// Filter
let filtered: Vec<_> = state
.tracked_requests
.iter()
.filter(|r| r.url.contains("example"))
.collect();
assert_eq!(filtered.len(), 1);
assert_eq!(filtered[0].url, "https://example.com");
// Clear
state.tracked_requests.clear();
assert!(state.tracked_requests.is_empty());
}
#[tokio::test]
async fn test_addscript_and_addinitscript_separate_dispatch() {
let mut state = DaemonState::new();
// Both should be handled (not "Not yet implemented") even without a browser
let cmd1 = json!({ "action": "addscript", "id": "as-1", "content": "console.log(1)" });
let result1 = execute_command(&cmd1, &mut state).await;
let err1 = result1["error"].as_str().unwrap_or("");
assert!(
!err1.contains("Not yet implemented"),
"addscript should be handled"
);
let cmd2 = json!({ "action": "addinitscript", "id": "ais-1", "script": "console.log(2)" });
let result2 = execute_command(&cmd2, &mut state).await;
let err2 = result2["error"].as_str().unwrap_or("");
assert!(
!err2.contains("Not yet implemented"),
"addinitscript should be handled"
);
}
#[tokio::test]
async fn test_frame_context_management() {
let mut state = DaemonState::new();
assert!(state.active_frame_id.is_none());
// Set a frame ID and verify it persists
state.active_frame_id = Some("child-frame-123".to_string());
assert_eq!(state.active_frame_id.as_deref(), Some("child-frame-123"));
// Clearing the frame ID (what mainframe does)
state.active_frame_id = None;
assert!(state.active_frame_id.is_none());
}
#[tokio::test]
async fn test_addstyle_supports_content_and_url() {
let mut state = DaemonState::new();
// Both content-based and url-based addstyle should be recognized
let cmd1 = json!({ "action": "addstyle", "id": "style-1", "content": "body { color: red }" });
let result1 = execute_command(&cmd1, &mut state).await;
let err1 = result1["error"].as_str().unwrap_or("");
assert!(!err1.contains("Not yet implemented"));
let cmd2 =
json!({ "action": "addstyle", "id": "style-2", "url": "https://example.com/style.css" });
let result2 = execute_command(&cmd2, &mut state).await;
let err2 = result2["error"].as_str().unwrap_or("");
assert!(!err2.contains("Not yet implemented"));
}
#[tokio::test]
async fn test_domain_filter_sanitize() {
use super::network::DomainFilter;
let filter = DomainFilter::new("example.com");
assert!(filter.is_allowed("example.com"));
assert!(!filter.is_allowed("evil.com"));
filter.check_url("https://example.com/path").unwrap();
assert!(filter.check_url("https://evil.com").is_err());
}
#[tokio::test]
async fn test_state_find_auto_returns_none_for_nonexistent() {
use super::state;
let result = state::find_auto_state_file("nonexistent-session-xyz");
assert!(result.is_none());
}
+217
View File
@@ -0,0 +1,217 @@
use serde::{Deserialize, Serialize};
use std::collections::HashSet;
use std::env;
use std::fs;
use std::path::PathBuf;
/// Result of a policy check for an action.
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum PolicyResult {
/// Action is allowed.
Allow,
/// Action is blocked with the given reason.
Deny(String),
/// Action requires confirmation before proceeding.
RequiresConfirmation,
}
/// Policy configuration loaded from a JSON file.
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ActionPolicy {
#[serde(skip)]
path: PathBuf,
#[serde(default)]
default: Option<String>,
#[serde(default)]
allow: Option<Vec<String>>,
#[serde(default)]
deny: Option<Vec<String>>,
#[serde(default)]
confirm: Option<Vec<String>>,
}
/// Confirmation categories parsed from AGENT_BROWSER_CONFIRM_ACTIONS.
#[derive(Debug, Clone)]
pub struct ConfirmActions {
pub categories: HashSet<String>,
}
impl ConfirmActions {
pub fn from_env() -> Option<Self> {
let val = env::var("AGENT_BROWSER_CONFIRM_ACTIONS").ok()?;
if val.is_empty() {
return None;
}
let categories: HashSet<String> = val
.split(',')
.map(|s| s.trim().to_lowercase())
.filter(|s| !s.is_empty())
.collect();
if categories.is_empty() {
None
} else {
Some(Self { categories })
}
}
pub fn requires_confirmation(&self, action: &str) -> bool {
self.categories.contains(action)
}
}
impl ActionPolicy {
/// Load policy from a JSON file at the given path.
pub fn load(path: &str) -> Result<Self, String> {
let path_buf = PathBuf::from(path);
let contents = fs::read_to_string(&path_buf)
.map_err(|e| format!("Failed to read policy file: {}", e))?;
let mut policy: ActionPolicy =
serde_json::from_str(&contents).map_err(|e| format!("Invalid policy JSON: {}", e))?;
policy.path = path_buf;
Ok(policy)
}
/// Load policy if AGENT_BROWSER_ACTION_POLICY env var is set.
/// Falls back to AGENT_BROWSER_POLICY for backwards compatibility.
pub fn load_if_exists() -> Option<Self> {
let path = env::var("AGENT_BROWSER_ACTION_POLICY")
.or_else(|_| env::var("AGENT_BROWSER_POLICY"))
.ok()?;
Self::load(&path).ok()
}
/// Check whether an action is allowed, denied, or requires confirmation.
pub fn check(&self, action: &str) -> PolicyResult {
if let Some(deny) = &self.deny {
if deny.iter().any(|a| a == action) {
return PolicyResult::Deny(format!("Action '{}' is denied by policy", action));
}
}
if let Some(confirm) = &self.confirm {
if confirm.iter().any(|a| a == action) {
return PolicyResult::RequiresConfirmation;
}
}
if let Some(allow) = &self.allow {
if !allow.is_empty() && !allow.iter().any(|a| a == action) {
let is_default_deny = self
.default
.as_deref()
.map(|d| d.eq_ignore_ascii_case("deny"))
.unwrap_or(true);
if is_default_deny {
return PolicyResult::Deny(format!(
"Action '{}' is not in the allow list",
action
));
}
}
} else if let Some(ref default) = self.default {
if default.eq_ignore_ascii_case("deny") {
return PolicyResult::Deny(format!(
"Action '{}' denied: default policy is deny",
action
));
}
}
PolicyResult::Allow
}
/// Reload policy from the file. Re-reads the JSON and updates the policy.
pub fn reload(&mut self) -> Result<(), String> {
let contents = fs::read_to_string(&self.path)
.map_err(|e| format!("Failed to read policy file: {}", e))?;
let mut policy: ActionPolicy =
serde_json::from_str(&contents).map_err(|e| format!("Invalid policy JSON: {}", e))?;
policy.path = self.path.clone();
*self = policy;
Ok(())
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::test_utils::EnvGuard;
#[test]
fn test_policy_allow_whitelist() {
let json = r#"{"allow": ["click", "type"], "deny": [], "confirm": []}"#;
let policy: ActionPolicy = serde_json::from_str(json).unwrap();
assert_eq!(policy.check("click"), PolicyResult::Allow);
assert_eq!(policy.check("type"), PolicyResult::Allow);
assert!(matches!(policy.check("navigate"), PolicyResult::Deny(_)));
}
#[test]
fn test_policy_deny() {
let json = r#"{"allow": [], "deny": ["delete"], "confirm": []}"#;
let policy: ActionPolicy = serde_json::from_str(json).unwrap();
assert!(matches!(policy.check("delete"), PolicyResult::Deny(_)));
}
#[test]
fn test_policy_confirm() {
let json = r#"{"allow": [], "deny": [], "confirm": ["submit"]}"#;
let policy: ActionPolicy = serde_json::from_str(json).unwrap();
assert_eq!(policy.check("submit"), PolicyResult::RequiresConfirmation);
}
#[test]
fn test_policy_deny_takes_precedence() {
let json = r#"{"allow": ["danger"], "deny": ["danger"], "confirm": []}"#;
let policy: ActionPolicy = serde_json::from_str(json).unwrap();
assert!(matches!(policy.check("danger"), PolicyResult::Deny(_)));
}
#[test]
fn test_policy_confirm_takes_precedence_over_allow() {
let json = r#"{"allow": ["submit"], "deny": [], "confirm": ["submit"]}"#;
let policy: ActionPolicy = serde_json::from_str(json).unwrap();
assert_eq!(policy.check("submit"), PolicyResult::RequiresConfirmation);
}
#[test]
fn test_policy_empty_allow_allows_all() {
let json = r#"{"allow": [], "deny": [], "confirm": []}"#;
let policy: ActionPolicy = serde_json::from_str(json).unwrap();
assert_eq!(policy.check("anything"), PolicyResult::Allow);
}
#[test]
fn test_policy_missing_allow_allows_all() {
let json = r#"{"deny": []}"#;
let policy: ActionPolicy = serde_json::from_str(json).unwrap();
assert_eq!(policy.check("anything"), PolicyResult::Allow);
}
#[test]
fn test_policy_default_allow() {
let json = r#"{"default": "allow", "deny": ["navigate"]}"#;
let policy: ActionPolicy = serde_json::from_str(json).unwrap();
assert_eq!(policy.check("click"), PolicyResult::Allow);
assert!(matches!(policy.check("navigate"), PolicyResult::Deny(_)));
}
#[test]
fn test_policy_default_deny() {
let json = r#"{"default": "deny", "allow": ["click"]}"#;
let policy: ActionPolicy = serde_json::from_str(json).unwrap();
assert_eq!(policy.check("click"), PolicyResult::Allow);
assert!(matches!(policy.check("navigate"), PolicyResult::Deny(_)));
}
#[test]
fn test_confirm_actions_from_env() {
let _guard = EnvGuard::new(&["AGENT_BROWSER_CONFIRM_ACTIONS"]);
_guard.set("AGENT_BROWSER_CONFIRM_ACTIONS", "navigate,click,fill");
let ca = ConfirmActions::from_env().unwrap();
assert!(ca.requires_confirmation("navigate"));
assert!(ca.requires_confirmation("click"));
assert!(ca.requires_confirmation("fill"));
assert!(!ca.requires_confirmation("screenshot"));
}
}
+274
View File
@@ -0,0 +1,274 @@
//! Browser provider connections for remote CDP sessions.
//!
//! Supports Browserbase, Browser Use, and Kernel providers. Each provider
//! returns a CDP WebSocket URL for connecting via BrowserManager.
use serde_json::{json, Value};
use std::env;
/// Provider session info for cleanup on failure.
pub struct ProviderSession {
pub provider: String,
pub session_id: String,
}
/// Connects to the specified browser provider and returns a CDP WebSocket URL
/// along with session info for cleanup on failure.
pub async fn connect_provider(
provider_name: &str,
) -> Result<(String, Option<ProviderSession>), String> {
match provider_name.to_lowercase().as_str() {
"browserbase" => connect_browserbase().await,
"browser-use" | "browseruse" => connect_browser_use().await,
"kernel" => connect_kernel().await,
_ => Err(format!(
"Unknown provider '{}'. Supported: browserbase, browser-use, kernel",
provider_name
)),
}
}
/// Close a provider session (call on CDP connect failure).
pub async fn close_provider_session(session: &ProviderSession) {
let client = reqwest::Client::new();
match session.provider.as_str() {
"browserbase" => {
if let Ok(api_key) = env::var("BROWSERBASE_API_KEY") {
let _ = client
.delete(format!(
"https://api.browserbase.com/v1/sessions/{}",
session.session_id
))
.header("X-BB-API-Key", &api_key)
.send()
.await;
}
}
"browser-use" => {
if let Ok(api_key) = env::var("BROWSER_USE_API_KEY") {
let _ = client
.patch(format!(
"https://api.browser-use.com/api/v2/browsers/{}",
session.session_id
))
.header("X-Browser-Use-API-Key", &api_key)
.header("Content-Type", "application/json")
.json(&json!({ "action": "stop" }))
.send()
.await;
}
}
"kernel" => {
if let Ok(api_key) = env::var("KERNEL_API_KEY") {
let endpoint = env::var("KERNEL_ENDPOINT")
.unwrap_or_else(|_| "https://api.onkernel.com".to_string());
let _ = client
.delete(format!(
"{}/browsers/{}",
endpoint.trim_end_matches('/'),
session.session_id
))
.header("Authorization", format!("Bearer {}", api_key))
.send()
.await;
}
}
_ => {}
}
}
async fn connect_browserbase() -> Result<(String, Option<ProviderSession>), String> {
let api_key = env::var("BROWSERBASE_API_KEY")
.map_err(|_| "BROWSERBASE_API_KEY environment variable is not set")?;
let project_id = env::var("BROWSERBASE_PROJECT_ID")
.map_err(|_| "BROWSERBASE_PROJECT_ID environment variable is not set")?;
let client = reqwest::Client::new();
let response = client
.post("https://api.browserbase.com/v1/sessions")
.header("Content-Type", "application/json")
.header("X-BB-API-Key", &api_key)
.json(&json!({ "projectId": project_id }))
.send()
.await
.map_err(|e| format!("Browserbase request failed: {}", e))?;
let status = response.status();
let body = response
.text()
.await
.map_err(|e| format!("Failed to read Browserbase response: {}", e))?;
if !status.is_success() {
return Err(format!(
"Browserbase API error ({}): {}",
status.as_u16(),
body
));
}
let json: Value =
serde_json::from_str(&body).map_err(|e| format!("Invalid Browserbase response: {}", e))?;
let session_id = json
.get("id")
.and_then(|v| v.as_str())
.unwrap_or("")
.to_string();
let ws_url = json
.get("connectUrl")
.and_then(|v| v.as_str())
.map(String::from)
.ok_or_else(|| "Browserbase response missing connectUrl".to_string())?;
Ok((
ws_url,
Some(ProviderSession {
provider: "browserbase".to_string(),
session_id,
}),
))
}
async fn connect_browser_use() -> Result<(String, Option<ProviderSession>), String> {
let api_key = env::var("BROWSER_USE_API_KEY")
.map_err(|_| "BROWSER_USE_API_KEY environment variable is not set")?;
let client = reqwest::Client::new();
let response = client
.post("https://api.browser-use.com/api/v2/browsers")
.header("Content-Type", "application/json")
.header("X-Browser-Use-API-Key", &api_key)
.json(&json!({}))
.send()
.await
.map_err(|e| format!("Browser Use request failed: {}", e))?;
let status = response.status();
let body = response
.text()
.await
.map_err(|e| format!("Failed to read Browser Use response: {}", e))?;
if !status.is_success() {
return Err(format!(
"Browser Use API error ({}): {}",
status.as_u16(),
body
));
}
let json: Value =
serde_json::from_str(&body).map_err(|e| format!("Invalid Browser Use response: {}", e))?;
let session_id = json
.get("id")
.and_then(|v| v.as_str())
.unwrap_or("")
.to_string();
let ws_url = json
.get("cdp_url")
.or_else(|| json.get("cdpUrl"))
.and_then(|v| v.as_str())
.map(String::from)
.ok_or_else(|| "Browser Use response missing cdp_url or cdpUrl".to_string())?;
Ok((
ws_url,
Some(ProviderSession {
provider: "browser-use".to_string(),
session_id,
}),
))
}
async fn connect_kernel() -> Result<(String, Option<ProviderSession>), String> {
let api_key =
env::var("KERNEL_API_KEY").map_err(|_| "KERNEL_API_KEY environment variable is not set")?;
let endpoint =
env::var("KERNEL_ENDPOINT").unwrap_or_else(|_| "https://api.onkernel.com".to_string());
let url = format!("{}/browsers", endpoint.trim_end_matches('/'));
let headless = env::var("KERNEL_HEADLESS")
.map(|v| v == "1" || v.eq_ignore_ascii_case("true"))
.unwrap_or(true);
let stealth = env::var("KERNEL_STEALTH")
.map(|v| v == "1" || v.eq_ignore_ascii_case("true"))
.unwrap_or(false);
let timeout_seconds = env::var("KERNEL_TIMEOUT_SECONDS")
.ok()
.and_then(|v| v.parse::<u64>().ok())
.unwrap_or(300);
let mut body = json!({
"headless": headless,
"stealth": stealth,
"timeout_seconds": timeout_seconds,
});
if let Ok(profile) = env::var("KERNEL_PROFILE_NAME") {
if !profile.is_empty() {
body.as_object_mut()
.unwrap()
.insert("profile".to_string(), json!(profile));
}
}
let client = reqwest::Client::new();
let response = client
.post(&url)
.header("Content-Type", "application/json")
.header("Authorization", format!("Bearer {}", api_key))
.json(&body)
.send()
.await
.map_err(|e| format!("Kernel request failed: {}", e))?;
let status = response.status();
let resp_body = response
.text()
.await
.map_err(|e| format!("Failed to read Kernel response: {}", e))?;
if !status.is_success() {
return Err(format!(
"Kernel API error ({}): {}",
status.as_u16(),
resp_body
));
}
let json: Value =
serde_json::from_str(&resp_body).map_err(|e| format!("Invalid Kernel response: {}", e))?;
let session_id = json
.get("session_id")
.or_else(|| json.get("id"))
.and_then(|v| v.as_str())
.unwrap_or("")
.to_string();
let ws_url = json
.get("cdp_ws_url")
.or_else(|| json.get("connectUrl"))
.or_else(|| json.get("connect_url"))
.or_else(|| json.get("cdpUrl"))
.or_else(|| json.get("cdp_url"))
.and_then(|v| v.as_str())
.map(String::from)
.ok_or_else(|| {
"Kernel response missing cdp_ws_url, connectUrl, connect_url, cdpUrl, or cdp_url"
.to_string()
})?;
Ok((
ws_url,
Some(ProviderSession {
provider: "kernel".to_string(),
session_id,
}),
))
}
+203
View File
@@ -0,0 +1,203 @@
use serde_json::{json, Value};
use std::path::PathBuf;
use std::process::Command;
pub struct RecordingState {
pub active: bool,
pub output_path: String,
pub temp_dir: PathBuf,
pub frame_count: u64,
}
impl RecordingState {
pub fn new() -> Self {
Self {
active: false,
output_path: String::new(),
temp_dir: PathBuf::new(),
frame_count: 0,
}
}
}
pub fn recording_start(state: &mut RecordingState, path: &str) -> Result<Value, String> {
if state.active {
return Err("Recording already active".to_string());
}
let timestamp = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap_or_default()
.as_millis();
let temp_dir = std::env::temp_dir().join(format!("agent-browser-recording-{}", timestamp));
let _ = std::fs::create_dir_all(&temp_dir);
state.active = true;
state.output_path = path.to_string();
state.temp_dir = temp_dir;
state.frame_count = 0;
Ok(json!({ "started": true, "path": path }))
}
pub fn recording_add_frame(state: &mut RecordingState, frame_data: &[u8]) {
if !state.active {
return;
}
let frame_path = state
.temp_dir
.join(format!("frame_{:06}.jpg", state.frame_count));
let _ = std::fs::write(&frame_path, frame_data);
state.frame_count += 1;
}
pub fn recording_stop(state: &mut RecordingState) -> Result<Value, String> {
if !state.active {
return Err("No recording in progress".to_string());
}
state.active = false;
if state.frame_count == 0 {
let _ = std::fs::remove_dir_all(&state.temp_dir);
return Err("No frames captured".to_string());
}
let frame_pattern = state
.temp_dir
.join("frame_%06d.jpg")
.to_string_lossy()
.to_string();
let output = &state.output_path;
// Encode with ffmpeg
let result = Command::new("ffmpeg")
.args([
"-y",
"-framerate",
"30",
"-i",
&frame_pattern,
"-c:v",
"libx264",
"-pix_fmt",
"yuv420p",
"-preset",
"fast",
output,
])
.output();
let _ = std::fs::remove_dir_all(&state.temp_dir);
match result {
Ok(output_result) => {
if output_result.status.success() {
Ok(json!({ "path": output, "frames": state.frame_count }))
} else {
let stderr = String::from_utf8_lossy(&output_result.stderr);
Err(format!(
"ffmpeg failed: {}",
stderr.chars().take(200).collect::<String>()
))
}
}
Err(e) => Err(format!(
"ffmpeg not found or failed to execute: {}. Install ffmpeg to enable recording.",
e
)),
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_recording_state_new() {
let state = RecordingState::new();
assert!(!state.active);
assert!(state.output_path.is_empty());
assert_eq!(state.frame_count, 0);
}
#[test]
fn test_recording_start_sets_active() {
let mut state = RecordingState::new();
let result = recording_start(&mut state, "/tmp/test.mp4");
assert!(result.is_ok());
assert!(state.active);
assert_eq!(state.output_path, "/tmp/test.mp4");
assert_eq!(state.frame_count, 0);
// Cleanup
let _ = std::fs::remove_dir_all(&state.temp_dir);
}
#[test]
fn test_recording_start_while_active() {
let mut state = RecordingState::new();
recording_start(&mut state, "/tmp/test1.mp4").unwrap();
let temp_dir = state.temp_dir.clone();
let result = recording_start(&mut state, "/tmp/test2.mp4");
assert!(result.is_err());
assert!(result.unwrap_err().contains("already active"));
let _ = std::fs::remove_dir_all(&temp_dir);
}
#[test]
fn test_recording_stop_not_active() {
let mut state = RecordingState::new();
let result = recording_stop(&mut state);
assert!(result.is_err());
assert!(result.unwrap_err().contains("No recording"));
}
#[test]
fn test_recording_stop_no_frames() {
let mut state = RecordingState::new();
recording_start(&mut state, "/tmp/test.mp4").unwrap();
let result = recording_stop(&mut state);
assert!(result.is_err());
assert!(result.unwrap_err().contains("No frames"));
assert!(!state.active);
}
#[test]
fn test_recording_add_frame_inactive() {
let mut state = RecordingState::new();
recording_add_frame(&mut state, b"fake-frame");
assert_eq!(state.frame_count, 0);
}
#[test]
fn test_recording_add_frame_active() {
let mut state = RecordingState::new();
recording_start(&mut state, "/tmp/test.mp4").unwrap();
recording_add_frame(&mut state, b"fake-frame-1");
recording_add_frame(&mut state, b"fake-frame-2");
assert_eq!(state.frame_count, 2);
let _ = std::fs::remove_dir_all(&state.temp_dir);
}
}
pub fn recording_restart(state: &mut RecordingState, path: &str) -> Result<Value, String> {
let previous = if state.active {
let stop_result = recording_stop(state);
stop_result
.ok()
.and_then(|v| v.get("path").and_then(|p| p.as_str()).map(String::from))
} else {
None
};
recording_start(state, path)?;
Ok(json!({
"restarted": true,
"previousPath": previous,
"path": path,
}))
}
+147
View File
@@ -0,0 +1,147 @@
use serde_json::Value;
use std::path::PathBuf;
use super::cdp::client::CdpClient;
use super::cdp::types::*;
use super::element::RefMap;
pub struct ScreenshotOptions {
pub selector: Option<String>,
pub path: Option<String>,
pub full_page: bool,
pub format: String,
pub quality: Option<i32>,
}
impl Default for ScreenshotOptions {
fn default() -> Self {
Self {
selector: None,
path: None,
full_page: false,
format: "png".to_string(),
quality: None,
}
}
}
pub async fn take_screenshot(
client: &CdpClient,
session_id: &str,
ref_map: &RefMap,
options: &ScreenshotOptions,
) -> Result<(String, String), String> {
let mut params = CaptureScreenshotParams {
format: Some(options.format.clone()),
quality: if options.format == "jpeg" {
options.quality.or(Some(80))
} else {
None
},
clip: None,
from_surface: Some(true),
capture_beyond_viewport: if options.full_page { Some(true) } else { None },
};
if options.full_page {
let metrics: Value = client
.send_command_no_params("Page.getLayoutMetrics", Some(session_id))
.await?;
let content_size = metrics
.get("contentSize")
.or_else(|| metrics.get("cssContentSize"));
if let Some(size) = content_size {
let width = size.get("width").and_then(|v| v.as_f64()).unwrap_or(1280.0);
let height = size.get("height").and_then(|v| v.as_f64()).unwrap_or(720.0);
params.clip = Some(Viewport {
x: 0.0,
y: 0.0,
width,
height,
scale: 1.0,
});
}
} else if let Some(ref selector) = options.selector {
// Element screenshot via bounding box
let object_id =
super::element::resolve_element_object_id(client, session_id, ref_map, selector)
.await?;
let result: EvaluateResult = client
.send_command_typed(
"Runtime.callFunctionOn",
&CallFunctionOnParams {
function_declaration: r#"function() {
const rect = this.getBoundingClientRect();
return { x: rect.x, y: rect.y, width: rect.width, height: rect.height };
}"#
.to_string(),
object_id: Some(object_id),
arguments: None,
return_by_value: Some(true),
await_promise: Some(false),
},
Some(session_id),
)
.await?;
if let Some(rect) = result.result.value {
let x = rect.get("x").and_then(|v| v.as_f64()).unwrap_or(0.0);
let y = rect.get("y").and_then(|v| v.as_f64()).unwrap_or(0.0);
let w = rect.get("width").and_then(|v| v.as_f64()).unwrap_or(100.0);
let h = rect.get("height").and_then(|v| v.as_f64()).unwrap_or(100.0);
params.clip = Some(Viewport {
x,
y,
width: w,
height: h,
scale: 1.0,
});
}
}
let result: CaptureScreenshotResult = client
.send_command_typed("Page.captureScreenshot", &params, Some(session_id))
.await?;
let ext = if options.format == "jpeg" {
"jpg"
} else {
"png"
};
let save_path = match &options.path {
Some(p) => p.clone(),
None => {
let dir = get_screenshot_dir();
let _ = std::fs::create_dir_all(&dir);
let timestamp = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap_or_default()
.as_millis();
let name = format!("screenshot-{}.{}", timestamp, ext);
dir.join(name).to_string_lossy().to_string()
}
};
let bytes = base64::Engine::decode(&base64::engine::general_purpose::STANDARD, &result.data)
.map_err(|e| format!("Failed to decode screenshot: {}", e))?;
std::fs::write(&save_path, &bytes)
.map_err(|e| format!("Failed to save screenshot to {}: {}", save_path, e))?;
Ok((save_path, result.data))
}
fn get_screenshot_dir() -> PathBuf {
if let Some(home) = dirs::home_dir() {
home.join(".agent-browser").join("tmp").join("screenshots")
} else {
std::env::temp_dir()
.join("agent-browser")
.join("screenshots")
}
}
+736
View File
@@ -0,0 +1,736 @@
use std::collections::HashMap;
use serde_json::Value;
use super::cdp::client::CdpClient;
use super::cdp::types::{
AXNode, AXProperty, AXValue, CallFunctionOnParams, EvaluateParams, EvaluateResult,
GetFullAXTreeResult,
};
use super::element::RefMap;
const INTERACTIVE_ROLES: &[&str] = &[
"button",
"link",
"textbox",
"checkbox",
"radio",
"combobox",
"listbox",
"menuitem",
"menuitemcheckbox",
"menuitemradio",
"option",
"searchbox",
"slider",
"spinbutton",
"switch",
"tab",
"treeitem",
];
const CONTENT_ROLES: &[&str] = &[
"heading",
"cell",
"gridcell",
"columnheader",
"rowheader",
"listitem",
"article",
"region",
"main",
"navigation",
];
const STRUCTURAL_ROLES: &[&str] = &[
"generic",
"group",
"list",
"table",
"row",
"rowgroup",
"grid",
"treegrid",
"menu",
"menubar",
"toolbar",
"tablist",
"tree",
"directory",
"document",
"application",
"presentation",
"none",
"WebArea",
"RootWebArea",
];
pub struct SnapshotOptions {
pub selector: Option<String>,
pub interactive: bool,
pub compact: bool,
pub depth: Option<usize>,
pub cursor: bool,
}
impl Default for SnapshotOptions {
fn default() -> Self {
Self {
selector: None,
interactive: false,
compact: false,
depth: None,
cursor: false,
}
}
}
struct TreeNode {
role: String,
name: String,
level: Option<i64>,
checked: Option<String>,
expanded: Option<bool>,
selected: Option<bool>,
disabled: Option<bool>,
required: Option<bool>,
value_text: Option<String>,
backend_node_id: Option<i64>,
children: Vec<usize>,
has_ref: bool,
ref_id: Option<String>,
depth: usize,
}
struct RoleNameTracker {
counts: HashMap<String, usize>,
entries: Vec<(usize, String)>,
}
impl RoleNameTracker {
fn new() -> Self {
Self {
counts: HashMap::new(),
entries: Vec::new(),
}
}
fn track(&mut self, role: &str, name: &str, node_idx: usize) -> usize {
let key = format!("{}:{}", role, name);
let count = self.counts.entry(key.clone()).or_insert(0);
let nth = *count;
*count += 1;
self.entries.push((node_idx, key));
nth
}
fn get_duplicates(&self) -> HashMap<String, usize> {
self.counts
.iter()
.filter(|(_, &count)| count > 1)
.map(|(key, &count)| (key.clone(), count))
.collect()
}
}
pub async fn take_snapshot(
client: &CdpClient,
session_id: &str,
options: &SnapshotOptions,
ref_map: &mut RefMap,
) -> Result<String, String> {
client
.send_command_no_params("DOM.enable", Some(session_id))
.await?;
client
.send_command_no_params("Accessibility.enable", Some(session_id))
.await?;
let ax_tree: GetFullAXTreeResult = client
.send_command_typed(
"Accessibility.getFullAXTree",
&serde_json::json!({}),
Some(session_id),
)
.await?;
let (tree_nodes, root_indices) = build_tree(&ax_tree.nodes);
let mut tracker = RoleNameTracker::new();
let mut next_ref: usize = ref_map.next_ref_num();
let mut nodes_with_refs: Vec<(usize, usize)> = Vec::new();
for (idx, node) in tree_nodes.iter().enumerate() {
let role = node.role.as_str();
let should_ref = if INTERACTIVE_ROLES.contains(&role) {
true
} else if CONTENT_ROLES.contains(&role) {
!node.name.is_empty()
} else {
false
};
if should_ref {
let nth = tracker.track(role, &node.name, idx);
nodes_with_refs.push((idx, nth));
}
}
let duplicates = tracker.get_duplicates();
let mut tree_nodes = tree_nodes;
for (idx, nth) in &nodes_with_refs {
let node = &tree_nodes[*idx];
let key = format!("{}:{}", node.role, node.name);
let actual_nth = if duplicates.contains_key(&key) {
Some(*nth)
} else {
None
};
let ref_id = format!("e{}", next_ref);
next_ref += 1;
ref_map.add(
ref_id.clone(),
tree_nodes[*idx].backend_node_id,
&tree_nodes[*idx].role,
&tree_nodes[*idx].name,
actual_nth,
);
tree_nodes[*idx].has_ref = true;
tree_nodes[*idx].ref_id = Some(ref_id);
}
ref_map.set_next_ref_num(next_ref);
let mut output = String::new();
for &root_idx in &root_indices {
render_tree(&tree_nodes, root_idx, 0, &mut output, options);
}
if options.compact {
output = compact_tree(&output, options.interactive);
}
let mut trimmed = output.trim().to_string();
if trimmed.is_empty() {
if options.interactive {
return Ok("(no interactive elements)".to_string());
}
return Ok("(empty page)".to_string());
}
if options.cursor {
let cursor_section = find_cursor_interactive_elements(client, session_id, ref_map).await?;
if !cursor_section.is_empty() {
trimmed.push_str("\n# Cursor-interactive elements:\n");
trimmed.push_str(&cursor_section);
}
}
Ok(trimmed)
}
async fn find_cursor_interactive_elements(
client: &CdpClient,
session_id: &str,
ref_map: &mut RefMap,
) -> Result<String, String> {
let js = r#"
(function() {
const elements = [];
const walker = document.createTreeWalker(document.body, NodeFilter.SHOW_ELEMENT);
let node;
while (node = walker.nextNode()) {
if (node.closest && node.closest('[hidden], [aria-hidden="true"]')) continue;
const explicitRole = node.getAttribute ? node.getAttribute('role') : null;
if (explicitRole) continue;
const tag = node.tagName ? node.tagName.toLowerCase() : '';
const hasClick = node.onclick || (node.attributes && node.attributes.getNamedItem('onclick'));
const tabindex = node.getAttribute ? node.getAttribute('tabindex') : null;
const contentEditable = node.getAttribute ? node.getAttribute('contenteditable') : null;
const isInherentlyClickable =
(tag === 'a' && node.href) || tag === 'button' ||
(tag === 'input' && ['submit','button','image','reset'].indexOf((node.type||'').toLowerCase()) >= 0) ||
tag === 'summary';
const isFocusable = tabindex !== null && parseInt(tabindex, 10) >= 0;
const isEditable = contentEditable === '' || contentEditable === 'true';
if (hasClick || isInherentlyClickable || isFocusable || isEditable) {
elements.push(node);
}
}
return elements;
})()
"#;
let result: EvaluateResult = client
.send_command_typed(
"Runtime.evaluate",
&EvaluateParams {
expression: js.to_string(),
return_by_value: Some(false),
await_promise: Some(false),
},
Some(session_id),
)
.await?;
let array_object_id = match result.result.object_id {
Some(id) => id,
None => return Ok(String::new()),
};
let props_result: Value = client
.send_command(
"Runtime.getProperties",
Some(serde_json::json!({ "objectId": array_object_id })),
Some(session_id),
)
.await?;
let empty: Vec<Value> = Vec::new();
let result_array = props_result
.get("result")
.and_then(|v| v.as_array())
.unwrap_or(&empty);
let mut indexed: Vec<(usize, String)> = Vec::new();
for prop in result_array {
let name = prop.get("name").and_then(|v| v.as_str()).unwrap_or("");
if let Ok(idx) = name.parse::<usize>() {
if let Some(obj_id) = prop
.get("value")
.and_then(|v| v.get("objectId"))
.and_then(|v| v.as_str())
{
indexed.push((idx, obj_id.to_string()));
}
}
}
indexed.sort_by_key(|(idx, _)| *idx);
let element_object_ids: Vec<String> = indexed.into_iter().map(|(_, id)| id).collect();
let mut next_ref = ref_map.next_ref_num();
let mut lines: Vec<String> = Vec::new();
let get_text_js =
r#"function(){ return (this.innerText || this.textContent || '').trim().slice(0, 100) }"#;
for object_id in &element_object_ids {
let describe: Value = client
.send_command(
"DOM.describeNode",
Some(serde_json::json!({ "objectId": object_id })),
Some(session_id),
)
.await?;
let backend_node_id = describe
.get("node")
.and_then(|n| n.get("backendNodeId"))
.and_then(|v| v.as_i64());
let text_result: EvaluateResult = client
.send_command_typed(
"Runtime.callFunctionOn",
&CallFunctionOnParams {
function_declaration: get_text_js.to_string(),
object_id: Some(object_id.clone()),
arguments: None,
return_by_value: Some(true),
await_promise: Some(false),
},
Some(session_id),
)
.await?;
let text = text_result
.result
.value
.as_ref()
.and_then(|v| v.as_str())
.unwrap_or("")
.trim()
.to_string();
let kind = "clickable";
let ref_id = format!("e{}", next_ref);
next_ref += 1;
ref_map.add(ref_id.clone(), backend_node_id, kind, &text, None);
let escaped = text
.replace('\\', "\\\\")
.replace('"', "\\\"")
.replace('\n', " ")
.replace('\r', " ");
lines.push(format!("[ref={}] ({}) \"{}\"", ref_id, kind, escaped));
}
ref_map.set_next_ref_num(next_ref);
Ok(lines.join("\n"))
}
fn build_tree(nodes: &[AXNode]) -> (Vec<TreeNode>, Vec<usize>) {
let mut tree_nodes: Vec<TreeNode> = Vec::with_capacity(nodes.len());
let mut id_to_idx: HashMap<String, usize> = HashMap::new();
for (i, node) in nodes.iter().enumerate() {
let role = extract_ax_string(&node.role);
let name = extract_ax_string(&node.name);
let value_text = extract_ax_string_opt(&node.value);
let (level, checked, expanded, selected, disabled, required) =
extract_properties(&node.properties);
if node.ignored.unwrap_or(false) && role != "RootWebArea" {
tree_nodes.push(TreeNode {
role: String::new(),
name: String::new(),
level: None,
checked: None,
expanded: None,
selected: None,
disabled: None,
required: None,
value_text: None,
backend_node_id: None,
children: Vec::new(),
has_ref: false,
ref_id: None,
depth: 0,
});
id_to_idx.insert(node.node_id.clone(), i);
continue;
}
tree_nodes.push(TreeNode {
role,
name,
level,
checked,
expanded,
selected,
disabled,
required,
value_text,
backend_node_id: node.backend_d_o_m_node_id,
children: Vec::new(),
has_ref: false,
ref_id: None,
depth: 0,
});
id_to_idx.insert(node.node_id.clone(), i);
}
// Build parent-child relationships
for (i, node) in nodes.iter().enumerate() {
if let Some(ref child_ids) = node.child_ids {
for cid in child_ids {
if let Some(&child_idx) = id_to_idx.get(cid) {
tree_nodes[i].children.push(child_idx);
}
}
}
}
// Set depths
let mut root_indices = Vec::new();
let children_exist: Vec<bool> = nodes.iter().map(|_| false).collect();
let mut is_child = children_exist;
for node in &tree_nodes {
for &child in &node.children {
is_child[child] = true;
}
}
for (i, &is_c) in is_child.iter().enumerate() {
if !is_c {
root_indices.push(i);
}
}
fn set_depth(nodes: &mut [TreeNode], idx: usize, depth: usize) {
nodes[idx].depth = depth;
let children: Vec<usize> = nodes[idx].children.clone();
for child_idx in children {
set_depth(nodes, child_idx, depth + 1);
}
}
for &root in &root_indices {
set_depth(&mut tree_nodes, root, 0);
}
(tree_nodes, root_indices)
}
fn render_tree(
nodes: &[TreeNode],
idx: usize,
indent: usize,
output: &mut String,
options: &SnapshotOptions,
) {
let node = &nodes[idx];
if node.role.is_empty() {
// Ignored node -- still render children
for &child in &node.children {
render_tree(nodes, child, indent, output, options);
}
return;
}
if let Some(max_depth) = options.depth {
if indent > max_depth {
return;
}
}
let role = &node.role;
// Skip root WebArea wrapper
if role == "RootWebArea" || role == "WebArea" {
for &child in &node.children {
render_tree(nodes, child, indent, output, options);
}
return;
}
if options.interactive && !node.has_ref {
// In interactive mode, skip non-interactive but render children
for &child in &node.children {
render_tree(nodes, child, indent, output, options);
}
return;
}
let prefix = " ".repeat(indent);
let mut line = format!("{}- {}", prefix, role);
if !node.name.is_empty() {
line.push_str(&format!(" \"{}\"", node.name));
}
// Properties
let mut attrs = Vec::new();
if let Some(level) = node.level {
attrs.push(format!("level={}", level));
}
if let Some(ref checked) = node.checked {
attrs.push(format!("checked={}", checked));
}
if let Some(expanded) = node.expanded {
attrs.push(format!("expanded={}", expanded));
}
if let Some(selected) = node.selected {
if selected {
attrs.push("selected".to_string());
}
}
if let Some(disabled) = node.disabled {
if disabled {
attrs.push("disabled".to_string());
}
}
if let Some(required) = node.required {
if required {
attrs.push("required".to_string());
}
}
if let Some(ref ref_id) = node.ref_id {
attrs.push(format!("ref={}", ref_id));
}
if !attrs.is_empty() {
line.push_str(&format!(" [{}]", attrs.join(", ")));
}
// Value
if let Some(ref val) = node.value_text {
if !val.is_empty() && val != &node.name {
line.push_str(&format!(": {}", val));
}
}
output.push_str(&line);
output.push('\n');
for &child in &node.children {
render_tree(nodes, child, indent + 1, output, options);
}
}
fn compact_tree(tree: &str, interactive: bool) -> String {
let lines: Vec<&str> = tree.lines().collect();
if lines.is_empty() {
return String::new();
}
let mut keep = vec![false; lines.len()];
for (i, line) in lines.iter().enumerate() {
if line.contains("[ref=") || line.contains(": ") {
keep[i] = true;
// Mark ancestors
let my_indent = count_indent(line);
for j in (0..i).rev() {
let ancestor_indent = count_indent(lines[j]);
if ancestor_indent < my_indent {
keep[j] = true;
if ancestor_indent == 0 {
break;
}
}
}
}
}
let result: Vec<&str> = lines
.iter()
.enumerate()
.filter(|(i, _)| keep[*i])
.map(|(_, line)| *line)
.collect();
let output = result.join("\n");
if output.trim().is_empty() && interactive {
return "(no interactive elements)".to_string();
}
output
}
fn count_indent(line: &str) -> usize {
let trimmed = line.trim_start();
(line.len() - trimmed.len()) / 2
}
fn extract_ax_string(value: &Option<AXValue>) -> String {
match value {
Some(v) => match &v.value {
Some(Value::String(s)) => s.clone(),
Some(Value::Number(n)) => n.to_string(),
Some(Value::Bool(b)) => b.to_string(),
_ => String::new(),
},
None => String::new(),
}
}
fn extract_ax_string_opt(value: &Option<AXValue>) -> Option<String> {
match value {
Some(v) => match &v.value {
Some(Value::String(s)) if !s.is_empty() => Some(s.clone()),
Some(Value::Number(n)) => Some(n.to_string()),
_ => None,
},
None => None,
}
}
type NodeProperties = (
Option<i64>, // level
Option<String>, // checked
Option<bool>, // expanded
Option<bool>, // selected
Option<bool>, // disabled
Option<bool>, // required
);
fn extract_properties(props: &Option<Vec<AXProperty>>) -> NodeProperties {
let mut level = None;
let mut checked = None;
let mut expanded = None;
let mut selected = None;
let mut disabled = None;
let mut required = None;
if let Some(properties) = props {
for prop in properties {
match prop.name.as_str() {
"level" => {
level = prop.value.value.as_ref().and_then(|v| v.as_i64());
}
"checked" => {
checked = prop.value.value.as_ref().map(|v| match v {
Value::String(s) => s.clone(),
Value::Bool(b) => b.to_string(),
_ => "false".to_string(),
});
}
"expanded" => {
expanded = prop.value.value.as_ref().and_then(|v| v.as_bool());
}
"selected" => {
selected = prop.value.value.as_ref().and_then(|v| v.as_bool());
}
"disabled" => {
disabled = prop.value.value.as_ref().and_then(|v| v.as_bool());
}
"required" => {
required = prop.value.value.as_ref().and_then(|v| v.as_bool());
}
_ => {}
}
}
}
(level, checked, expanded, selected, disabled, required)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_interactive_roles() {
assert!(INTERACTIVE_ROLES.contains(&"button"));
assert!(INTERACTIVE_ROLES.contains(&"textbox"));
assert!(!INTERACTIVE_ROLES.contains(&"heading"));
}
#[test]
fn test_content_roles() {
assert!(CONTENT_ROLES.contains(&"heading"));
assert!(!CONTENT_ROLES.contains(&"button"));
}
#[test]
fn test_compact_tree_basic() {
let tree = "- navigation\n - link \"Home\" [ref=e1]\n - link \"About\" [ref=e2]\n- main\n - heading \"Title\"\n - paragraph\n - text: Hello\n";
let result = compact_tree(tree, false);
assert!(result.contains("[ref=e1]"));
assert!(result.contains("[ref=e2]"));
assert!(result.contains("Hello"));
}
#[test]
fn test_compact_tree_empty_interactive() {
let result = compact_tree("- generic\n", true);
assert_eq!(result, "(no interactive elements)");
}
#[test]
fn test_count_indent() {
assert_eq!(count_indent("- heading"), 0);
assert_eq!(count_indent(" - link"), 1);
assert_eq!(count_indent(" - text"), 2);
}
#[test]
fn test_role_name_tracker() {
let mut tracker = RoleNameTracker::new();
assert_eq!(tracker.track("button", "Submit", 0), 0);
assert_eq!(tracker.track("button", "Submit", 1), 1);
assert_eq!(tracker.track("button", "Cancel", 2), 0);
let dups = tracker.get_duplicates();
assert!(dups.contains_key("button:Submit"));
assert!(!dups.contains_key("button:Cancel"));
}
}
+607
View File
@@ -0,0 +1,607 @@
use aes_gcm::{aead::Aead, aead::KeyInit, Aes256Gcm};
use serde::{Deserialize, Serialize};
use serde_json::{json, Value};
use sha2::{Digest, Sha256};
use std::fs;
use std::path::PathBuf;
use super::cdp::client::CdpClient;
use super::cdp::types::EvaluateParams;
use super::cookies::{self, Cookie};
#[derive(Debug, Serialize, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct StorageState {
pub cookies: Vec<Cookie>,
pub origins: Vec<OriginStorage>,
}
#[derive(Debug, Serialize, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct OriginStorage {
pub origin: String,
pub local_storage: Vec<StorageEntry>,
#[serde(default)]
pub session_storage: Vec<StorageEntry>,
}
#[derive(Debug, Serialize, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct StorageEntry {
pub name: String,
pub value: String,
}
pub async fn save_state(
client: &CdpClient,
session_id: &str,
path: Option<&str>,
session_name: Option<&str>,
session_id_str: &str,
) -> Result<String, String> {
let cookies = cookies::get_cookies(client, session_id, None).await?;
// Get current origin's storage
let origin_js = r#"(() => {
const result = { origin: location.origin, localStorage: [], sessionStorage: [] };
try {
for (let i = 0; i < localStorage.length; i++) {
const key = localStorage.key(i);
result.localStorage.push({ name: key, value: localStorage.getItem(key) });
}
} catch(e) {}
try {
for (let i = 0; i < sessionStorage.length; i++) {
const key = sessionStorage.key(i);
result.sessionStorage.push({ name: key, value: sessionStorage.getItem(key) });
}
} catch(e) {}
return result;
})()"#;
let origin_result: super::cdp::types::EvaluateResult = client
.send_command_typed(
"Runtime.evaluate",
&EvaluateParams {
expression: origin_js.to_string(),
return_by_value: Some(true),
await_promise: Some(false),
},
Some(session_id),
)
.await?;
let origin_data = origin_result.result.value.unwrap_or(Value::Null);
let origins = if origin_data.is_object() {
let origin = origin_data
.get("origin")
.and_then(|v| v.as_str())
.unwrap_or("")
.to_string();
let local_storage: Vec<StorageEntry> = origin_data
.get("localStorage")
.and_then(|v| serde_json::from_value(v.clone()).ok())
.unwrap_or_default();
let session_storage: Vec<StorageEntry> = origin_data
.get("sessionStorage")
.and_then(|v| serde_json::from_value(v.clone()).ok())
.unwrap_or_default();
if !origin.is_empty() && origin != "null" {
vec![OriginStorage {
origin,
local_storage,
session_storage,
}]
} else {
vec![]
}
} else {
vec![]
};
let state = StorageState { cookies, origins };
let json_str = serde_json::to_string_pretty(&state)
.map_err(|e| format!("Failed to serialize state: {}", e))?;
let mut save_path = match path {
Some(p) => p.to_string(),
None => {
let dir = get_sessions_dir();
let _ = fs::create_dir_all(&dir);
let name = session_name.unwrap_or("default");
dir.join(format!("{}-{}.json", name, session_id_str))
.to_string_lossy()
.to_string()
}
};
if let Ok(key) = std::env::var("AGENT_BROWSER_ENCRYPTION_KEY") {
let encrypted = encrypt_data(json_str.as_bytes(), &key)?;
save_path.push_str(".enc");
fs::write(&save_path, &encrypted)
.map_err(|e| format!("Failed to write state to {}: {}", save_path, e))?;
} else {
fs::write(&save_path, &json_str)
.map_err(|e| format!("Failed to write state to {}: {}", save_path, e))?;
}
Ok(save_path)
}
pub async fn load_state(client: &CdpClient, session_id: &str, path: &str) -> Result<(), String> {
let json_str = if path.ends_with(".enc") {
let key = std::env::var("AGENT_BROWSER_ENCRYPTION_KEY").map_err(|_| {
"Encrypted state file requires AGENT_BROWSER_ENCRYPTION_KEY".to_string()
})?;
let data =
fs::read(path).map_err(|e| format!("Failed to read state from {}: {}", path, e))?;
let decrypted = decrypt_data(&data, &key)?;
String::from_utf8(decrypted)
.map_err(|e| format!("Decrypted state is not valid UTF-8: {}", e))?
} else {
match fs::read_to_string(path) {
Ok(s) => s,
Err(e) => {
if let Ok(key) = std::env::var("AGENT_BROWSER_ENCRYPTION_KEY") {
let enc_path = format!("{}.enc", path);
if let Ok(data) = fs::read(&enc_path) {
let decrypted = decrypt_data(&data, &key)?;
String::from_utf8(decrypted)
.map_err(|de| format!("Decrypted state is not valid UTF-8: {}", de))?
} else {
return Err(format!("Failed to read state from {}: {}", path, e));
}
} else {
return Err(format!("Failed to read state from {}: {}", path, e));
}
}
}
};
let state: StorageState =
serde_json::from_str(&json_str).map_err(|e| format!("Invalid state file: {}", e))?;
// Load cookies
if !state.cookies.is_empty() {
let cookie_values: Vec<Value> = state
.cookies
.iter()
.map(|c| serde_json::to_value(c).unwrap_or(Value::Null))
.collect();
cookies::set_cookies(client, session_id, cookie_values, None).await?;
}
// Load storage per origin
for origin in &state.origins {
if origin.local_storage.is_empty() && origin.session_storage.is_empty() {
continue;
}
// Navigate to origin to set storage
let navigate_url = format!("{}/", origin.origin.trim_end_matches('/'));
client
.send_command(
"Page.navigate",
Some(json!({ "url": navigate_url })),
Some(session_id),
)
.await?;
// Brief wait for navigation
tokio::time::sleep(tokio::time::Duration::from_millis(500)).await;
for entry in &origin.local_storage {
let js = format!(
"localStorage.setItem({}, {})",
serde_json::to_string(&entry.name).unwrap_or_default(),
serde_json::to_string(&entry.value).unwrap_or_default(),
);
let _ = client
.send_command_typed::<_, super::cdp::types::EvaluateResult>(
"Runtime.evaluate",
&EvaluateParams {
expression: js,
return_by_value: Some(true),
await_promise: Some(false),
},
Some(session_id),
)
.await;
}
for entry in &origin.session_storage {
let js = format!(
"sessionStorage.setItem({}, {})",
serde_json::to_string(&entry.name).unwrap_or_default(),
serde_json::to_string(&entry.value).unwrap_or_default(),
);
let _ = client
.send_command_typed::<_, super::cdp::types::EvaluateResult>(
"Runtime.evaluate",
&EvaluateParams {
expression: js,
return_by_value: Some(true),
await_promise: Some(false),
},
Some(session_id),
)
.await;
}
}
Ok(())
}
fn is_state_file(path: &std::path::Path) -> bool {
let fname = path
.file_name()
.unwrap_or_default()
.to_string_lossy()
.to_string();
fname.ends_with(".json") || fname.ends_with(".json.enc")
}
fn is_encrypted_state(path: &std::path::Path) -> bool {
path.to_string_lossy().ends_with(".json.enc")
}
pub fn state_list() -> Result<Value, String> {
let dir = get_sessions_dir();
if !dir.exists() {
return Ok(json!({ "files": [], "directory": dir.to_string_lossy() }));
}
let mut files = Vec::new();
let entries = fs::read_dir(&dir).map_err(|e| format!("Failed to read sessions dir: {}", e))?;
for entry in entries.flatten() {
let path = entry.path();
if is_state_file(&path) {
let metadata = fs::metadata(&path).ok();
let filename = path
.file_name()
.unwrap_or_default()
.to_string_lossy()
.to_string();
let size = metadata.as_ref().map(|m| m.len()).unwrap_or(0);
let modified = metadata
.as_ref()
.and_then(|m| m.modified().ok())
.and_then(|t| t.duration_since(std::time::UNIX_EPOCH).ok())
.map(|d| d.as_secs())
.unwrap_or(0);
let encrypted = is_encrypted_state(&path);
files.push(json!({
"filename": filename,
"path": path.to_string_lossy(),
"size": size,
"modified": modified,
"encrypted": encrypted,
}));
}
}
Ok(json!({ "files": files, "directory": dir.to_string_lossy() }))
}
pub fn state_show(path: &str) -> Result<Value, String> {
let encrypted = path.ends_with(".enc");
let json_str = if encrypted {
let key = std::env::var("AGENT_BROWSER_ENCRYPTION_KEY").map_err(|_| {
"Encrypted state file requires AGENT_BROWSER_ENCRYPTION_KEY".to_string()
})?;
let data = fs::read(path).map_err(|e| format!("Failed to read state file: {}", e))?;
let decrypted = decrypt_data(&data, &key)?;
String::from_utf8(decrypted)
.map_err(|e| format!("Decrypted state is not valid UTF-8: {}", e))?
} else {
fs::read_to_string(path).map_err(|e| format!("Failed to read state file: {}", e))?
};
let state: StorageState =
serde_json::from_str(&json_str).map_err(|e| format!("Invalid state file: {}", e))?;
let metadata = fs::metadata(path).ok();
let filename = std::path::Path::new(path)
.file_name()
.unwrap_or_default()
.to_string_lossy()
.to_string();
Ok(json!({
"filename": filename,
"path": path,
"size": metadata.as_ref().map(|m| m.len()).unwrap_or(0),
"modified": metadata.as_ref()
.and_then(|m| m.modified().ok())
.and_then(|t| t.duration_since(std::time::UNIX_EPOCH).ok())
.map(|d| d.as_secs())
.unwrap_or(0),
"encrypted": encrypted,
"summary": format!("{} cookies, {} origins", state.cookies.len(), state.origins.len()),
"state": state,
}))
}
pub fn state_clear(path: Option<&str>) -> Result<Value, String> {
if let Some(p) = path {
fs::remove_file(p).map_err(|e| format!("Failed to delete state: {}", e))?;
return Ok(json!({ "deleted": p }));
}
let dir = get_sessions_dir();
if !dir.exists() {
return Ok(json!({ "deleted": 0 }));
}
let mut count = 0;
if let Ok(entries) = fs::read_dir(&dir) {
for entry in entries.flatten() {
let path = entry.path();
if is_state_file(&path) {
let _ = fs::remove_file(&path);
count += 1;
}
}
}
Ok(json!({ "deleted": count }))
}
pub fn state_clean(max_age_days: u64) -> Result<Value, String> {
let dir = get_sessions_dir();
if !dir.exists() {
return Ok(json!({ "cleaned": 0, "keptCount": 0, "days": max_age_days }));
}
let now = std::time::SystemTime::now();
let max_age = std::time::Duration::from_secs(max_age_days * 86400);
let mut deleted = 0;
let mut kept = 0;
if let Ok(entries) = fs::read_dir(&dir) {
for entry in entries.flatten() {
let path = entry.path();
if !is_state_file(&path) {
continue;
}
if let Ok(metadata) = fs::metadata(&path) {
if let Ok(modified) = metadata.modified() {
if let Ok(age) = now.duration_since(modified) {
if age > max_age {
let _ = fs::remove_file(&path);
deleted += 1;
continue;
}
}
}
}
kept += 1;
}
}
Ok(json!({ "cleaned": deleted, "keptCount": kept, "days": max_age_days }))
}
pub fn state_rename(old_path: &str, new_name: &str) -> Result<Value, String> {
let old = PathBuf::from(old_path);
if !old.exists() {
return Err(format!("State file not found: {}", old_path));
}
let fallback = PathBuf::from(".");
let dir = old.parent().unwrap_or(&fallback);
let new_path = dir.join(format!("{}.json", new_name));
fs::rename(&old, &new_path).map_err(|e| format!("Failed to rename state: {}", e))?;
Ok(json!({
"renamed": true,
"from": old_path,
"to": new_path.to_string_lossy(),
}))
}
fn encrypt_data(data: &[u8], key_str: &str) -> Result<Vec<u8>, String> {
let mut hasher = Sha256::new();
hasher.update(key_str.as_bytes());
let key_bytes = hasher.finalize();
let cipher =
Aes256Gcm::new_from_slice(&key_bytes).map_err(|e| format!("Invalid key: {}", e))?;
let mut nonce = [0u8; 12];
getrandom::getrandom(&mut nonce).map_err(|e| format!("Failed to generate nonce: {}", e))?;
let ciphertext = cipher
.encrypt(aes_gcm::Nonce::from_slice(&nonce), data)
.map_err(|e| format!("Encryption failed: {}", e))?;
let mut result = Vec::with_capacity(12 + ciphertext.len());
result.extend_from_slice(&nonce);
result.extend_from_slice(&ciphertext);
Ok(result)
}
fn decrypt_data(data: &[u8], key_str: &str) -> Result<Vec<u8>, String> {
if data.len() < 13 {
return Err("Ciphertext too short".to_string());
}
let (nonce_bytes, ciphertext) = data.split_at(12);
let mut hasher = Sha256::new();
hasher.update(key_str.as_bytes());
let key_bytes = hasher.finalize();
let cipher =
Aes256Gcm::new_from_slice(&key_bytes).map_err(|e| format!("Invalid key: {}", e))?;
let plaintext = cipher
.decrypt(aes_gcm::Nonce::from_slice(nonce_bytes), ciphertext)
.map_err(|e| format!("Decryption failed: {}", e))?;
Ok(plaintext)
}
pub fn find_auto_state_file(session_name: &str) -> Option<String> {
let dir = get_sessions_dir();
if !dir.exists() {
return None;
}
let prefix = format!("{}-", session_name);
let mut best_path: Option<(String, std::time::SystemTime)> = None;
if let Ok(entries) = fs::read_dir(&dir) {
for entry in entries.flatten() {
let path = entry.path();
let fname = path
.file_name()
.unwrap_or_default()
.to_string_lossy()
.to_string();
let is_match = fname.starts_with(&prefix)
&& (fname.ends_with(".json") || fname.ends_with(".json.enc"));
if !is_match {
continue;
}
let modified = fs::metadata(&path)
.ok()
.and_then(|m| m.modified().ok())
.unwrap_or(std::time::UNIX_EPOCH);
if best_path.as_ref().map_or(true, |(_, t)| modified > *t) {
best_path = Some((path.to_string_lossy().to_string(), modified));
}
}
}
best_path.map(|(p, _)| p)
}
pub fn get_sessions_dir() -> PathBuf {
if let Some(home) = dirs::home_dir() {
home.join(".agent-browser").join("sessions")
} else {
std::env::temp_dir().join("agent-browser").join("sessions")
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_storage_state_serialization() {
let state = StorageState {
cookies: vec![Cookie {
name: "session".to_string(),
value: "abc123".to_string(),
domain: ".example.com".to_string(),
path: "/".to_string(),
expires: 0.0,
size: 0,
http_only: true,
secure: false,
session: true,
same_site: Some("Lax".to_string()),
}],
origins: vec![OriginStorage {
origin: "https://example.com".to_string(),
local_storage: vec![StorageEntry {
name: "key".to_string(),
value: "val".to_string(),
}],
session_storage: vec![],
}],
};
let json = serde_json::to_string_pretty(&state).unwrap();
let parsed: StorageState = serde_json::from_str(&json).unwrap();
assert_eq!(parsed.cookies.len(), 1);
assert_eq!(parsed.cookies[0].name, "session");
assert_eq!(parsed.origins.len(), 1);
assert_eq!(parsed.origins[0].local_storage.len(), 1);
}
#[test]
fn test_storage_state_empty() {
let state = StorageState {
cookies: vec![],
origins: vec![],
};
let json = serde_json::to_string(&state).unwrap();
let parsed: StorageState = serde_json::from_str(&json).unwrap();
assert!(parsed.cookies.is_empty());
assert!(parsed.origins.is_empty());
}
#[test]
fn test_state_show_nonexistent_file() {
let result = state_show("/tmp/nonexistent-agent-browser-state-file.json");
assert!(result.is_err());
}
#[test]
fn test_state_clear_nonexistent_file() {
let result = state_clear(Some("/tmp/nonexistent-agent-browser-state-file.json"));
assert!(result.is_err());
}
#[test]
fn test_state_rename_nonexistent() {
let result = state_rename("/tmp/nonexistent-agent-browser-state-file.json", "new-name");
assert!(result.is_err());
assert!(result.unwrap_err().contains("not found"));
}
#[test]
fn test_state_list_returns_json() {
let result = state_list().unwrap();
assert!(result.get("files").is_some());
assert!(result.get("directory").is_some());
}
#[test]
fn test_sessions_dir_path() {
let dir = get_sessions_dir();
assert!(dir.to_string_lossy().contains("sessions"));
}
#[test]
fn test_encrypt_decrypt_roundtrip() {
let plain = b"hello world";
let key = "test-secret-key";
let encrypted = encrypt_data(plain, key).unwrap();
assert!(encrypted.len() > 12);
assert_ne!(&encrypted[12..], plain);
let decrypted = decrypt_data(&encrypted, key).unwrap();
assert_eq!(decrypted, plain);
}
#[test]
fn test_decrypt_wrong_key_fails() {
let plain = b"secret data";
let encrypted = encrypt_data(plain, "key1").unwrap();
let result = decrypt_data(&encrypted, "key2");
assert!(result.is_err());
}
#[test]
fn test_cookie_serde_roundtrip() {
let cookie = Cookie {
name: "test".to_string(),
value: "123".to_string(),
domain: ".test.com".to_string(),
path: "/api".to_string(),
expires: 1700000000.0,
size: 7,
http_only: false,
secure: true,
session: false,
same_site: Some("Strict".to_string()),
};
let json = serde_json::to_value(&cookie).unwrap();
assert_eq!(json["name"], "test");
assert_eq!(json["httpOnly"], false);
assert_eq!(json["secure"], true);
assert_eq!(json["sameSite"], "Strict");
}
}
+94
View File
@@ -0,0 +1,94 @@
use serde_json::{json, Value};
use super::cdp::client::CdpClient;
use super::cdp::types::EvaluateParams;
pub async fn storage_get(
client: &CdpClient,
session_id: &str,
storage_type: &str,
key: Option<&str>,
) -> Result<Value, String> {
let st = storage_js_name(storage_type);
if let Some(k) = key {
let js = format!(
"{}.getItem({})",
st,
serde_json::to_string(k).unwrap_or_default()
);
let result = eval_simple(client, session_id, &js).await?;
Ok(json!({ "key": k, "value": result }))
} else {
let js = format!(
r#"(() => {{
const s = {};
const data = {{}};
for (let i = 0; i < s.length; i++) {{
const key = s.key(i);
data[key] = s.getItem(key);
}}
return data;
}})()"#,
st
);
let result = eval_simple(client, session_id, &js).await?;
Ok(json!({ "data": result }))
}
}
pub async fn storage_set(
client: &CdpClient,
session_id: &str,
storage_type: &str,
key: &str,
value: &str,
) -> Result<(), String> {
let st = storage_js_name(storage_type);
let js = format!(
"{}.setItem({}, {})",
st,
serde_json::to_string(key).unwrap_or_default(),
serde_json::to_string(value).unwrap_or_default(),
);
eval_simple(client, session_id, &js).await?;
Ok(())
}
pub async fn storage_clear(
client: &CdpClient,
session_id: &str,
storage_type: &str,
) -> Result<(), String> {
let st = storage_js_name(storage_type);
let js = format!("{}.clear()", st);
eval_simple(client, session_id, &js).await?;
Ok(())
}
fn storage_js_name(storage_type: &str) -> &str {
match storage_type {
"session" => "sessionStorage",
_ => "localStorage",
}
}
async fn eval_simple(client: &CdpClient, session_id: &str, js: &str) -> Result<Value, String> {
let result: super::cdp::types::EvaluateResult = client
.send_command_typed(
"Runtime.evaluate",
&EvaluateParams {
expression: js.to_string(),
return_by_value: Some(true),
await_promise: Some(false),
},
Some(session_id),
)
.await?;
if let Some(ref details) = result.exception_details {
return Err(format!("Storage error: {}", details.text));
}
Ok(result.result.value.unwrap_or(Value::Null))
}
+385
View File
@@ -0,0 +1,385 @@
use serde_json::{json, Value};
use std::net::SocketAddr;
use std::sync::Arc;
use futures_util::{SinkExt, StreamExt};
use tokio::net::TcpListener;
use tokio::sync::{broadcast, Mutex};
use tokio_tungstenite::tungstenite::Message;
use super::cdp::client::CdpClient;
/// Frame metadata from CDP Page.screencastFrame events.
#[derive(Debug, Clone)]
pub struct FrameMetadata {
pub offset_top: f64,
pub page_scale_factor: f64,
pub device_width: u32,
pub device_height: u32,
pub scroll_offset_x: f64,
pub scroll_offset_y: f64,
pub timestamp: u64,
}
impl Default for FrameMetadata {
fn default() -> Self {
Self {
offset_top: 0.0,
page_scale_factor: 1.0,
device_width: 1280,
device_height: 720,
scroll_offset_x: 0.0,
scroll_offset_y: 0.0,
timestamp: 0,
}
}
}
pub struct StreamServer {
port: u16,
frame_tx: broadcast::Sender<String>,
client_count: Arc<Mutex<usize>>,
}
impl StreamServer {
pub async fn start(
preferred_port: u16,
client: Arc<CdpClient>,
session_id: String,
) -> Result<Self, String> {
let addr = format!("127.0.0.1:{}", preferred_port);
let listener = TcpListener::bind(&addr)
.await
.map_err(|e| format!("Failed to bind stream server: {}", e))?;
let actual_addr = listener
.local_addr()
.map_err(|e| format!("Failed to get stream address: {}", e))?;
let port = actual_addr.port();
let (frame_tx, _) = broadcast::channel::<String>(64);
let client_count = Arc::new(Mutex::new(0usize));
let frame_tx_clone = frame_tx.clone();
let client_count_clone = client_count.clone();
tokio::spawn(async move {
accept_loop(
listener,
frame_tx_clone,
client_count_clone,
client,
session_id,
)
.await;
});
Ok(Self {
port,
frame_tx,
client_count,
})
}
pub fn port(&self) -> u16 {
self.port
}
/// Broadcast a raw frame string (legacy).
pub fn broadcast_frame(&self, frame_json: &str) {
let _ = self.frame_tx.send(frame_json.to_string());
}
/// Broadcast a screencast frame with structured metadata.
pub fn broadcast_screencast_frame(&self, base64_data: &str, metadata: &FrameMetadata) {
let msg = json!({
"type": "frame",
"data": base64_data,
"metadata": {
"offsetTop": metadata.offset_top,
"pageScaleFactor": metadata.page_scale_factor,
"deviceWidth": metadata.device_width,
"deviceHeight": metadata.device_height,
"scrollOffsetX": metadata.scroll_offset_x,
"scrollOffsetY": metadata.scroll_offset_y,
"timestamp": metadata.timestamp,
}
});
let _ = self.frame_tx.send(msg.to_string());
}
/// Broadcast a status message to all connected clients.
pub fn broadcast_status(
&self,
connected: bool,
screencasting: bool,
viewport_width: u32,
viewport_height: u32,
) {
let msg = json!({
"type": "status",
"connected": connected,
"screencasting": screencasting,
"viewportWidth": viewport_width,
"viewportHeight": viewport_height,
});
let _ = self.frame_tx.send(msg.to_string());
}
/// Broadcast an error message to all connected clients.
pub fn broadcast_error(&self, message: &str) {
let msg = json!({
"type": "error",
"message": message,
});
let _ = self.frame_tx.send(msg.to_string());
}
}
async fn accept_loop(
listener: TcpListener,
frame_tx: broadcast::Sender<String>,
client_count: Arc<Mutex<usize>>,
cdp_client: Arc<CdpClient>,
session_id: String,
) {
while let Ok((stream, addr)) = listener.accept().await {
let frame_rx = frame_tx.subscribe();
let client_count = client_count.clone();
let cdp = cdp_client.clone();
let sid = session_id.clone();
tokio::spawn(async move {
handle_ws_client(stream, addr, frame_rx, client_count, cdp, sid).await;
});
}
}
async fn handle_ws_client(
stream: tokio::net::TcpStream,
_addr: SocketAddr,
mut frame_rx: broadcast::Receiver<String>,
client_count: Arc<Mutex<usize>>,
cdp_client: Arc<CdpClient>,
session_id: String,
) {
// Origin checking on WebSocket handshake
let callback =
|req: &tokio_tungstenite::tungstenite::handshake::server::Request,
resp: tokio_tungstenite::tungstenite::handshake::server::Response| {
let origin = req
.headers()
.get("origin")
.and_then(|v| v.to_str().ok())
.map(|s| s.to_string());
if !is_allowed_origin(origin.as_deref()) {
let mut reject =
tokio_tungstenite::tungstenite::handshake::server::ErrorResponse::new(Some(
"Origin not allowed".to_string(),
));
*reject.status_mut() = tokio_tungstenite::tungstenite::http::StatusCode::FORBIDDEN;
return Err(reject);
}
Ok(resp)
};
let ws_stream = match tokio_tungstenite::accept_hdr_async(stream, callback).await {
Ok(ws) => ws,
Err(_) => return,
};
{
let mut count = client_count.lock().await;
*count += 1;
}
let (mut ws_tx, mut ws_rx) = ws_stream.split();
loop {
tokio::select! {
frame = frame_rx.recv() => {
match frame {
Ok(data) => {
if ws_tx.send(Message::Text(data)).await.is_err() {
break;
}
}
Err(_) => break,
}
}
msg = ws_rx.next() => {
match msg {
Some(Ok(Message::Text(text))) => {
handle_client_message(&text, &cdp_client, &session_id).await;
}
Some(Ok(Message::Close(_))) | None => break,
_ => {}
}
}
}
}
{
let mut count = client_count.lock().await;
*count = count.saturating_sub(1);
}
}
async fn handle_client_message(msg: &str, client: &CdpClient, session_id: &str) {
let parsed: Value = match serde_json::from_str(msg) {
Ok(v) => v,
Err(_) => return,
};
let msg_type = parsed.get("type").and_then(|v| v.as_str()).unwrap_or("");
match msg_type {
"input_mouse" => {
let _ = client
.send_command(
"Input.dispatchMouseEvent",
Some(json!({
"type": parsed.get("eventType").and_then(|v| v.as_str()).unwrap_or("mouseMoved"),
"x": parsed.get("x").and_then(|v| v.as_f64()).unwrap_or(0.0),
"y": parsed.get("y").and_then(|v| v.as_f64()).unwrap_or(0.0),
"button": parsed.get("button").and_then(|v| v.as_str()).unwrap_or("none"),
"clickCount": parsed.get("clickCount").and_then(|v| v.as_i64()).unwrap_or(0),
"deltaX": parsed.get("deltaX").and_then(|v| v.as_f64()).unwrap_or(0.0),
"deltaY": parsed.get("deltaY").and_then(|v| v.as_f64()).unwrap_or(0.0),
"modifiers": parsed.get("modifiers").and_then(|v| v.as_i64()).unwrap_or(0),
})),
Some(session_id),
)
.await;
}
"input_keyboard" => {
let _ = client
.send_command(
"Input.dispatchKeyEvent",
Some(json!({
"type": parsed.get("eventType").and_then(|v| v.as_str()).unwrap_or("keyDown"),
"key": parsed.get("key"),
"code": parsed.get("code"),
"text": parsed.get("text"),
"modifiers": parsed.get("modifiers").and_then(|v| v.as_i64()).unwrap_or(0),
})),
Some(session_id),
)
.await;
}
"input_touch" => {
let _ = client
.send_command(
"Input.dispatchTouchEvent",
Some(json!({
"type": parsed.get("eventType").and_then(|v| v.as_str()).unwrap_or("touchStart"),
"touchPoints": parsed.get("touchPoints").unwrap_or(&json!([])),
"modifiers": parsed.get("modifiers").and_then(|v| v.as_i64()).unwrap_or(0),
})),
Some(session_id),
)
.await;
}
"status" => {
// Client requesting status -- handled via broadcast_status from the caller
}
_ => {}
}
}
pub fn is_allowed_origin(origin: Option<&str>) -> bool {
match origin {
None => true,
Some(o) => {
if o.starts_with("file://") {
return true;
}
if let Ok(url) = url::Url::parse(o) {
let host = url.host_str().unwrap_or("");
host == "localhost" || host == "127.0.0.1" || host == "::1" || host == "[::1]"
} else {
false
}
}
}
}
pub async fn start_screencast(
client: &CdpClient,
session_id: &str,
format: &str,
quality: i32,
max_width: i32,
max_height: i32,
) -> Result<(), String> {
client
.send_command(
"Page.startScreencast",
Some(json!({
"format": format,
"quality": quality,
"maxWidth": max_width,
"maxHeight": max_height,
"everyNthFrame": 1,
})),
Some(session_id),
)
.await?;
Ok(())
}
pub async fn stop_screencast(client: &CdpClient, session_id: &str) -> Result<(), String> {
client
.send_command_no_params("Page.stopScreencast", Some(session_id))
.await?;
Ok(())
}
pub async fn ack_screencast_frame(
client: &CdpClient,
session_id: &str,
screencast_session_id: i64,
) -> Result<(), String> {
client
.send_command(
"Page.screencastFrameAck",
Some(json!({ "sessionId": screencast_session_id })),
Some(session_id),
)
.await?;
Ok(())
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_allowed_origin_none() {
assert!(is_allowed_origin(None));
}
#[test]
fn test_allowed_origin_file() {
assert!(is_allowed_origin(Some("file:///path/to/file")));
}
#[test]
fn test_allowed_origin_localhost() {
assert!(is_allowed_origin(Some("http://localhost:3000")));
assert!(is_allowed_origin(Some("http://127.0.0.1:8080")));
}
#[test]
fn test_disallowed_origin() {
assert!(!is_allowed_origin(Some("http://evil.com")));
}
#[test]
fn test_frame_metadata_default() {
let meta = FrameMetadata::default();
assert_eq!(meta.device_width, 1280);
assert_eq!(meta.device_height, 720);
assert_eq!(meta.page_scale_factor, 1.0);
}
}
+373
View File
@@ -0,0 +1,373 @@
use serde_json::{json, Value};
use std::path::PathBuf;
use super::cdp::client::CdpClient;
const MAX_PROFILE_EVENTS: usize = 5_000_000;
const DEFAULT_PROFILER_CATEGORIES: &[&str] = &[
"devtools.timeline",
"disabled-by-default-devtools.timeline",
"disabled-by-default-devtools.timeline.frame",
"disabled-by-default-devtools.timeline.stack",
"v8.execute",
"disabled-by-default-v8.cpu_profiler",
"disabled-by-default-v8.cpu_profiler.hires",
"v8",
"disabled-by-default-v8.runtime_stats",
"blink",
"blink.user_timing",
"latencyInfo",
"renderer.scheduler",
"sequence_manager",
"toplevel",
];
pub struct TracingState {
pub active: bool,
pub events: Vec<Value>,
pub events_dropped: bool,
}
impl TracingState {
pub fn new() -> Self {
Self {
active: false,
events: Vec::new(),
events_dropped: false,
}
}
}
pub async fn trace_start(
client: &CdpClient,
session_id: &str,
tracing_state: &mut TracingState,
) -> Result<Value, String> {
if tracing_state.active {
return Err("Tracing already active".to_string());
}
client
.send_command(
"Tracing.start",
Some(json!({
"traceConfig": {
"recordMode": "recordContinuously",
},
"transferMode": "ReturnAsStream",
})),
Some(session_id),
)
.await?;
tracing_state.active = true;
tracing_state.events.clear();
tracing_state.events_dropped = false;
Ok(json!({ "started": true }))
}
pub async fn trace_stop(
client: &CdpClient,
session_id: &str,
tracing_state: &mut TracingState,
path: Option<&str>,
) -> Result<Value, String> {
if !tracing_state.active {
return Err("No tracing in progress".to_string());
}
// Subscribe to events before stopping
let mut rx = client.subscribe();
client
.send_command_no_params("Tracing.end", Some(session_id))
.await?;
// Collect trace data with timeout
let mut trace_events: Vec<Value> = Vec::new();
let mut stream_handle: Option<String> = None;
let deadline = tokio::time::Instant::now() + tokio::time::Duration::from_secs(30);
loop {
let result = tokio::time::timeout_at(deadline, rx.recv()).await;
match result {
Ok(Ok(event)) => {
if event.session_id.as_deref() != Some(session_id) {
continue;
}
match event.method.as_str() {
"Tracing.dataCollected" => {
if let Some(arr) = event.params.get("value").and_then(|v| v.as_array()) {
trace_events.extend(arr.iter().cloned());
}
}
"Tracing.tracingComplete" => {
stream_handle = event
.params
.get("stream")
.and_then(|v| v.as_str())
.map(String::from);
break;
}
_ => {}
}
}
Ok(Err(_)) => break,
Err(_) => {
return Err("Tracing stop timed out after 30s".to_string());
}
}
}
// If ReturnAsStream mode was used, read trace data from the IO stream
if let Some(handle) = stream_handle {
if trace_events.is_empty() {
let stream_data = read_io_stream(client, session_id, &handle).await?;
if let Ok(parsed) = serde_json::from_str::<Value>(&stream_data) {
if let Some(events) = parsed.get("traceEvents").and_then(|v| v.as_array()) {
trace_events.extend(events.iter().cloned());
}
} else {
// Try parsing as newline-delimited JSON
for line in stream_data.lines() {
if let Ok(val) = serde_json::from_str::<Value>(line) {
if let Some(events) = val.get("traceEvents").and_then(|v| v.as_array()) {
trace_events.extend(events.iter().cloned());
} else {
trace_events.push(val);
}
}
}
}
}
// Close the IO stream
let _ = client
.send_command(
"IO.close",
Some(json!({ "handle": handle })),
Some(session_id),
)
.await;
}
tracing_state.active = false;
let save_path = match path {
Some(p) => p.to_string(),
None => {
let dir = get_traces_dir();
let _ = std::fs::create_dir_all(&dir);
let timestamp = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap_or_default()
.as_millis();
dir.join(format!("trace-{}.json", timestamp))
.to_string_lossy()
.to_string()
}
};
let trace_json = json!({ "traceEvents": trace_events });
let json_str = serde_json::to_string(&trace_json)
.map_err(|e| format!("Failed to serialize trace: {}", e))?;
std::fs::write(&save_path, json_str)
.map_err(|e| format!("Failed to write trace to {}: {}", save_path, e))?;
Ok(json!({ "path": save_path, "eventCount": trace_events.len() }))
}
pub async fn profiler_start(
client: &CdpClient,
session_id: &str,
tracing_state: &mut TracingState,
categories: Option<Vec<String>>,
) -> Result<Value, String> {
if tracing_state.active {
return Err("Profiling/tracing already active".to_string());
}
let cats: Vec<String> = categories.unwrap_or_else(|| {
DEFAULT_PROFILER_CATEGORIES
.iter()
.map(|s| s.to_string())
.collect()
});
client
.send_command(
"Tracing.start",
Some(json!({
"traceConfig": {
"includedCategories": cats,
"enableSampling": true,
},
"transferMode": "ReportEvents",
})),
Some(session_id),
)
.await?;
tracing_state.active = true;
tracing_state.events.clear();
tracing_state.events_dropped = false;
Ok(json!({ "started": true }))
}
pub async fn profiler_stop(
client: &CdpClient,
session_id: &str,
tracing_state: &mut TracingState,
path: Option<&str>,
) -> Result<Value, String> {
if !tracing_state.active {
return Err("No profiling in progress".to_string());
}
let mut rx = client.subscribe();
client
.send_command_no_params("Tracing.end", Some(session_id))
.await?;
let mut events: Vec<Value> = Vec::new();
let mut dropped = false;
let deadline = tokio::time::Instant::now() + tokio::time::Duration::from_secs(30);
loop {
let result = tokio::time::timeout_at(deadline, rx.recv()).await;
match result {
Ok(Ok(event)) => {
if event.session_id.as_deref() != Some(session_id) {
continue;
}
match event.method.as_str() {
"Tracing.dataCollected" => {
if let Some(arr) = event.params.get("value").and_then(|v| v.as_array()) {
if events.len() + arr.len() > MAX_PROFILE_EVENTS {
dropped = true;
} else {
events.extend(arr.iter().cloned());
}
}
}
"Tracing.tracingComplete" => {
break;
}
_ => {}
}
}
Ok(Err(_)) => break,
Err(_) => {
return Err("Profiler stop timed out after 30s".to_string());
}
}
}
tracing_state.active = false;
let save_path = match path {
Some(p) => p.to_string(),
None => {
let dir = get_profiles_dir();
let _ = std::fs::create_dir_all(&dir);
let timestamp = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap_or_default()
.as_millis();
dir.join(format!("profile-{}.json", timestamp))
.to_string_lossy()
.to_string()
}
};
let clock_domain = get_clock_domain();
let mut profile = json!({ "traceEvents": events });
if let Some(cd) = clock_domain {
profile
.as_object_mut()
.unwrap()
.insert("metadata".to_string(), json!({ "clock-domain": cd }));
}
let json_str = serde_json::to_string(&profile)
.map_err(|e| format!("Failed to serialize profile: {}", e))?;
std::fs::write(&save_path, json_str)
.map_err(|e| format!("Failed to write profile to {}: {}", save_path, e))?;
let event_count = events.len();
let mut result = json!({ "path": save_path, "eventCount": event_count });
if dropped {
result.as_object_mut().unwrap().insert(
"warning".to_string(),
Value::String(format!(
"Events exceeded {} limit; some dropped",
MAX_PROFILE_EVENTS
)),
);
}
Ok(result)
}
/// Read all data from a CDP IO stream handle.
async fn read_io_stream(
client: &CdpClient,
session_id: &str,
handle: &str,
) -> Result<String, String> {
let mut data = String::new();
loop {
let result = client
.send_command(
"IO.read",
Some(json!({
"handle": handle,
"size": 1024 * 1024,
})),
Some(session_id),
)
.await?;
if let Some(chunk) = result.get("data").and_then(|v| v.as_str()) {
data.push_str(chunk);
}
let eof = result.get("eof").and_then(|v| v.as_bool()).unwrap_or(true);
if eof {
break;
}
}
Ok(data)
}
fn get_clock_domain() -> Option<&'static str> {
if cfg!(target_os = "linux") {
Some("LINUX_CLOCK_MONOTONIC")
} else if cfg!(target_os = "macos") {
Some("MAC_MACH_ABSOLUTE_TIME")
} else {
None
}
}
fn get_traces_dir() -> PathBuf {
if let Some(home) = dirs::home_dir() {
home.join(".agent-browser").join("tmp").join("traces")
} else {
std::env::temp_dir().join("agent-browser").join("traces")
}
}
fn get_profiles_dir() -> PathBuf {
if let Some(home) = dirs::home_dir() {
home.join(".agent-browser").join("tmp").join("profiles")
} else {
std::env::temp_dir().join("agent-browser").join("profiles")
}
}
+201
View File
@@ -0,0 +1,201 @@
use serde_json::{json, Value};
use std::process::{Child, Command, Stdio};
use std::time::Duration;
use super::client::WebDriverClient;
const APPIUM_DEFAULT_PORT: u16 = 4723;
const APPIUM_STARTUP_TIMEOUT_SECS: u64 = 30;
pub struct AppiumManager {
pub client: WebDriverClient,
appium_process: Option<Child>,
pub device_udid: Option<String>,
}
impl AppiumManager {
pub async fn connect_or_launch(device_udid: Option<&str>) -> Result<Self, String> {
let port = APPIUM_DEFAULT_PORT;
let client = WebDriverClient::new(port);
// Check if Appium is already running
if is_appium_running(port).await {
return Ok(Self {
client,
appium_process: None,
device_udid: device_udid.map(String::from),
});
}
// Try to launch Appium
let appium_process = launch_appium(port)?;
// Wait for Appium to be ready
wait_for_appium(port, APPIUM_STARTUP_TIMEOUT_SECS).await?;
Ok(Self {
client,
appium_process: Some(appium_process),
device_udid: device_udid.map(String::from),
})
}
pub async fn create_ios_session(
&mut self,
device_name: Option<&str>,
platform_version: Option<&str>,
) -> Result<Value, String> {
let mut caps = json!({
"platformName": "iOS",
"automationName": "XCUITest",
"browserName": "Safari",
"noReset": true,
});
if let Some(name) = device_name {
caps["deviceName"] = json!(name);
} else {
caps["deviceName"] = json!("iPhone");
}
if let Some(ver) = platform_version {
caps["platformVersion"] = json!(ver);
}
if let Some(ref udid) = self.device_udid {
caps["udid"] = json!(udid);
}
self.client.create_session(caps).await
}
pub async fn tap(&self, x: f64, y: f64) -> Result<(), String> {
let sid = self
.client
.session_id_pub()
.ok_or("No active session")?
.to_string();
let actions = json!({
"actions": [{
"type": "pointer",
"id": "finger1",
"parameters": { "pointerType": "touch" },
"actions": [
{ "type": "pointerMove", "duration": 0, "x": x as i64, "y": y as i64 },
{ "type": "pointerDown", "button": 0 },
{ "type": "pause", "duration": 100 },
{ "type": "pointerUp", "button": 0 },
]
}]
});
self.client.execute_actions(&sid, &actions).await
}
pub async fn swipe(
&self,
start_x: f64,
start_y: f64,
end_x: f64,
end_y: f64,
duration_ms: u64,
) -> Result<(), String> {
let sid = self
.client
.session_id_pub()
.ok_or("No active session")?
.to_string();
let actions = json!({
"actions": [{
"type": "pointer",
"id": "finger1",
"parameters": { "pointerType": "touch" },
"actions": [
{ "type": "pointerMove", "duration": 0, "x": start_x as i64, "y": start_y as i64 },
{ "type": "pointerDown", "button": 0 },
{ "type": "pointerMove", "duration": duration_ms, "x": end_x as i64, "y": end_y as i64 },
{ "type": "pointerUp", "button": 0 },
]
}]
});
self.client.execute_actions(&sid, &actions).await
}
pub async fn close(&mut self) -> Result<(), String> {
let _ = self.client.delete_session().await;
if let Some(ref mut child) = self.appium_process {
let _ = child.kill();
let _ = child.wait();
}
Ok(())
}
}
impl Drop for AppiumManager {
fn drop(&mut self) {
if let Some(ref mut child) = self.appium_process {
let _ = child.kill();
let _ = child.wait();
}
}
}
async fn is_appium_running(port: u16) -> bool {
let addr = format!("127.0.0.1:{}", port);
tokio::time::timeout(
Duration::from_secs(2),
tokio::net::TcpStream::connect(&addr),
)
.await
.map(|r| r.is_ok())
.unwrap_or(false)
}
fn launch_appium(port: u16) -> Result<Child, String> {
// Try npx appium first, then direct appium
let result = Command::new("npx")
.args(["appium", "--relaxed-security", "--port", &port.to_string()])
.stdin(Stdio::null())
.stdout(Stdio::piped())
.stderr(Stdio::piped())
.spawn();
match result {
Ok(child) => Ok(child),
Err(_) => Command::new("appium")
.args(["--relaxed-security", "--port", &port.to_string()])
.stdin(Stdio::null())
.stdout(Stdio::piped())
.stderr(Stdio::piped())
.spawn()
.map_err(|e| {
format!(
"Failed to launch Appium. Install it with: npm install -g appium. Error: {}",
e
)
}),
}
}
async fn wait_for_appium(port: u16, timeout_secs: u64) -> Result<(), String> {
let deadline = tokio::time::Instant::now() + Duration::from_secs(timeout_secs);
loop {
if tokio::time::Instant::now() > deadline {
return Err("Timeout waiting for Appium to start".to_string());
}
if is_appium_running(port).await {
return Ok(());
}
tokio::time::sleep(Duration::from_millis(500)).await;
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_appium_constants() {
assert_eq!(APPIUM_DEFAULT_PORT, 4723);
assert_eq!(APPIUM_STARTUP_TIMEOUT_SECS, 30);
}
}
+142
View File
@@ -0,0 +1,142 @@
use async_trait::async_trait;
use serde_json::Value;
/// Abstract backend for browser automation. CDP (Chromium) and WebDriver
/// (Safari/iOS) share this interface so actions.rs can remain backend-agnostic
/// in the future.
#[async_trait]
pub trait BrowserBackend: Send + Sync {
async fn navigate(&self, url: &str) -> Result<(), String>;
async fn get_url(&self) -> Result<String, String>;
async fn get_title(&self) -> Result<String, String>;
async fn get_content(&self) -> Result<String, String>;
async fn evaluate(&self, script: &str) -> Result<Value, String>;
async fn screenshot(&self) -> Result<String, String>;
async fn click(&self, selector: &str) -> Result<(), String>;
async fn fill(&self, selector: &str, value: &str) -> Result<(), String>;
async fn close(&mut self) -> Result<(), String>;
async fn back(&self) -> Result<(), String>;
async fn forward(&self) -> Result<(), String>;
async fn reload(&self) -> Result<(), String>;
async fn get_cookies(&self) -> Result<Value, String>;
fn backend_type(&self) -> &str;
fn supports(&self, feature: &str) -> bool {
match feature {
"navigate" | "evaluate" | "screenshot" | "click" | "fill" => true,
"screencast" | "tracing" | "network_intercept" | "cdp" => self.backend_type() == "cdp",
_ => false,
}
}
fn unsupported_error(&self, action: &str) -> String {
format!(
"Action '{}' is not supported on the {} backend",
action,
self.backend_type()
)
}
}
/// WebDriver implementation of BrowserBackend
pub struct WebDriverBackend {
client: super::client::WebDriverClient,
}
impl WebDriverBackend {
pub fn new(client: super::client::WebDriverClient) -> Self {
Self { client }
}
}
#[async_trait]
impl BrowserBackend for WebDriverBackend {
async fn navigate(&self, url: &str) -> Result<(), String> {
self.client.navigate(url).await
}
async fn get_url(&self) -> Result<String, String> {
self.client.get_url().await
}
async fn get_title(&self) -> Result<String, String> {
self.client.get_title().await
}
async fn get_content(&self) -> Result<String, String> {
self.client.get_page_source().await
}
async fn evaluate(&self, script: &str) -> Result<Value, String> {
self.client.execute_script(script, vec![]).await
}
async fn screenshot(&self) -> Result<String, String> {
self.client.screenshot().await
}
async fn click(&self, selector: &str) -> Result<(), String> {
let element_id = self.client.find_element("css selector", selector).await?;
self.client.click_element(&element_id).await
}
async fn fill(&self, selector: &str, value: &str) -> Result<(), String> {
let element_id = self.client.find_element("css selector", selector).await?;
self.client.clear_element(&element_id).await?;
self.client.send_keys(&element_id, value).await
}
async fn close(&mut self) -> Result<(), String> {
self.client.delete_session().await
}
async fn back(&self) -> Result<(), String> {
self.client.back().await
}
async fn forward(&self) -> Result<(), String> {
self.client.forward().await
}
async fn reload(&self) -> Result<(), String> {
self.client.refresh().await
}
async fn get_cookies(&self) -> Result<Value, String> {
self.client.get_cookies().await
}
fn backend_type(&self) -> &str {
"webdriver"
}
}
/// CDP-backed backend constants for unsupported actions on WebDriver
pub const WEBDRIVER_UNSUPPORTED_ACTIONS: &[&str] = &[
"screencast_start",
"screencast_stop",
"trace_start",
"trace_stop",
"profiler_start",
"profiler_stop",
"route",
"unroute",
"expose",
"addscript",
"addinitscript",
"network",
"har_start",
"har_stop",
];
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_unsupported_actions() {
assert!(WEBDRIVER_UNSUPPORTED_ACTIONS.contains(&"screencast_start"));
assert!(WEBDRIVER_UNSUPPORTED_ACTIONS.contains(&"trace_start"));
assert!(!WEBDRIVER_UNSUPPORTED_ACTIONS.contains(&"navigate"));
}
}
+318
View File
@@ -0,0 +1,318 @@
use serde_json::{json, Value};
use std::time::Duration;
pub struct WebDriverClient {
base_url: String,
session_id: Option<String>,
}
impl WebDriverClient {
pub fn new(port: u16) -> Self {
Self {
base_url: format!("http://127.0.0.1:{}", port),
session_id: None,
}
}
pub async fn create_session(&mut self, capabilities: Value) -> Result<Value, String> {
let body = json!({
"capabilities": {
"alwaysMatch": capabilities,
}
});
let response = self.post("/session", &body).await?;
let session_id = response
.get("value")
.and_then(|v| v.get("sessionId"))
.and_then(|v| v.as_str())
.ok_or("No sessionId in response")?
.to_string();
self.session_id = Some(session_id);
Ok(response)
}
pub async fn delete_session(&mut self) -> Result<(), String> {
if let Some(ref sid) = self.session_id.clone() {
let _ = self.delete(&format!("/session/{}", sid)).await;
self.session_id = None;
}
Ok(())
}
pub async fn navigate(&self, url: &str) -> Result<(), String> {
let sid = self.session_id()?.to_string();
self.post(&format!("/session/{}/url", sid), &json!({ "url": url }))
.await?;
Ok(())
}
pub async fn get_url(&self) -> Result<String, String> {
let sid = self.session_id()?.to_string();
let response = self.get(&format!("/session/{}/url", sid)).await?;
Ok(response
.get("value")
.and_then(|v| v.as_str())
.unwrap_or("")
.to_string())
}
pub async fn get_title(&self) -> Result<String, String> {
let sid = self.session_id()?.to_string();
let response = self.get(&format!("/session/{}/title", sid)).await?;
Ok(response
.get("value")
.and_then(|v| v.as_str())
.unwrap_or("")
.to_string())
}
pub async fn find_element(&self, using: &str, value: &str) -> Result<String, String> {
let sid = self.session_id()?.to_string();
let response = self
.post(
&format!("/session/{}/element", sid),
&json!({ "using": using, "value": value }),
)
.await?;
let element_value = response.get("value").ok_or("No element in response")?;
element_value
.get("element-6066-11e4-a52e-4f735466cecf")
.or_else(|| element_value.get("ELEMENT"))
.and_then(|v| v.as_str())
.map(String::from)
.ok_or("No element ID in response".to_string())
}
pub async fn click_element(&self, element_id: &str) -> Result<(), String> {
let sid = self.session_id()?.to_string();
self.post(
&format!("/session/{}/element/{}/click", sid, element_id),
&json!({}),
)
.await?;
Ok(())
}
pub async fn send_keys(&self, element_id: &str, text: &str) -> Result<(), String> {
let sid = self.session_id()?.to_string();
self.post(
&format!("/session/{}/element/{}/value", sid, element_id),
&json!({ "text": text }),
)
.await?;
Ok(())
}
pub async fn clear_element(&self, element_id: &str) -> Result<(), String> {
let sid = self.session_id()?.to_string();
self.post(
&format!("/session/{}/element/{}/clear", sid, element_id),
&json!({}),
)
.await?;
Ok(())
}
pub async fn execute_script(&self, script: &str, args: Vec<Value>) -> Result<Value, String> {
let sid = self.session_id()?.to_string();
let response = self
.post(
&format!("/session/{}/execute/sync", sid),
&json!({ "script": script, "args": args }),
)
.await?;
Ok(response.get("value").cloned().unwrap_or(Value::Null))
}
pub async fn screenshot(&self) -> Result<String, String> {
let sid = self.session_id()?.to_string();
let response = self.get(&format!("/session/{}/screenshot", sid)).await?;
Ok(response
.get("value")
.and_then(|v| v.as_str())
.unwrap_or("")
.to_string())
}
pub async fn get_cookies(&self) -> Result<Value, String> {
let sid = self.session_id()?.to_string();
let response = self.get(&format!("/session/{}/cookie", sid)).await?;
Ok(response.get("value").cloned().unwrap_or(Value::Null))
}
pub async fn get_page_source(&self) -> Result<String, String> {
let sid = self.session_id()?.to_string();
let response = self.get(&format!("/session/{}/source", sid)).await?;
Ok(response
.get("value")
.and_then(|v| v.as_str())
.unwrap_or("")
.to_string())
}
pub async fn back(&self) -> Result<(), String> {
let sid = self.session_id()?.to_string();
self.post(&format!("/session/{}/back", sid), &json!({}))
.await?;
Ok(())
}
pub async fn forward(&self) -> Result<(), String> {
let sid = self.session_id()?.to_string();
self.post(&format!("/session/{}/forward", sid), &json!({}))
.await?;
Ok(())
}
pub async fn refresh(&self) -> Result<(), String> {
let sid = self.session_id()?.to_string();
self.post(&format!("/session/{}/refresh", sid), &json!({}))
.await?;
Ok(())
}
pub fn session_id_pub(&self) -> Option<&str> {
self.session_id.as_deref()
}
pub fn new_with_session(port: u16, session_id: String) -> Self {
Self {
base_url: format!("http://127.0.0.1:{}", port),
session_id: Some(session_id),
}
}
pub async fn execute_actions(&self, session_id: &str, actions: &Value) -> Result<(), String> {
self.post(&format!("/session/{}/actions", session_id), actions)
.await?;
Ok(())
}
fn session_id(&self) -> Result<&str, String> {
self.session_id
.as_deref()
.ok_or("No active WebDriver session".to_string())
}
async fn get(&self, path: &str) -> Result<Value, String> {
http_request("GET", &format!("{}{}", self.base_url, path), None).await
}
async fn post(&self, path: &str, body: &Value) -> Result<Value, String> {
http_request("POST", &format!("{}{}", self.base_url, path), Some(body)).await
}
async fn delete(&self, path: &str) -> Result<Value, String> {
http_request("DELETE", &format!("{}{}", self.base_url, path), None).await
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_client_new() {
let client = WebDriverClient::new(4444);
assert_eq!(client.base_url, "http://127.0.0.1:4444");
assert!(client.session_id.is_none());
}
#[test]
fn test_session_id_none() {
let client = WebDriverClient::new(4444);
let result = client.session_id();
assert!(result.is_err());
assert!(result.unwrap_err().contains("No active WebDriver session"));
}
#[test]
fn test_client_custom_port() {
let client = WebDriverClient::new(9515);
assert_eq!(client.base_url, "http://127.0.0.1:9515");
}
}
async fn http_request(method: &str, url: &str, body: Option<&Value>) -> Result<Value, String> {
let parsed = url::Url::parse(url).map_err(|e| format!("Invalid URL: {}", e))?;
let host = parsed.host_str().unwrap_or("127.0.0.1");
let port = parsed.port().unwrap_or(80);
let path = parsed.path();
let addr = format!("{}:{}", host, port);
let stream = tokio::time::timeout(
Duration::from_secs(10),
tokio::net::TcpStream::connect(&addr),
)
.await
.map_err(|_| format!("Connection timeout: {}", addr))?
.map_err(|e| format!("Connection failed: {}", e))?;
use tokio::io::{AsyncReadExt, AsyncWriteExt};
let body_str = body
.map(|b| serde_json::to_string(b).unwrap_or_default())
.unwrap_or_default();
let request = if body.is_some() {
format!(
"{} {} HTTP/1.1\r\nHost: {}\r\nContent-Type: application/json\r\nContent-Length: {}\r\nConnection: close\r\n\r\n{}",
method, path, addr, body_str.len(), body_str
)
} else {
format!(
"{} {} HTTP/1.1\r\nHost: {}\r\nConnection: close\r\n\r\n",
method, path, addr
)
};
let mut stream = stream;
stream
.write_all(request.as_bytes())
.await
.map_err(|e| format!("Write failed: {}", e))?;
let mut response = Vec::new();
stream
.read_to_end(&mut response)
.await
.map_err(|e| format!("Read failed: {}", e))?;
let response_str = String::from_utf8_lossy(&response);
let body_part = response_str.split("\r\n\r\n").nth(1).unwrap_or("").trim();
// Handle chunked encoding
let json_body = if body_part.contains('\n')
&& body_part
.chars()
.next()
.map(|c| c.is_ascii_hexdigit())
.unwrap_or(false)
{
// Chunked: skip chunk size lines
body_part
.lines()
.filter(|l| !l.chars().all(|c| c.is_ascii_hexdigit() || c == '\r'))
.collect::<Vec<&str>>()
.join("")
} else {
body_part.to_string()
};
if json_body.is_empty() {
return Ok(json!({}));
}
serde_json::from_str(&json_body).map_err(|e| {
format!(
"Invalid JSON response: {} (body: {})",
e,
json_body.chars().take(100).collect::<String>()
)
})
}
+235
View File
@@ -0,0 +1,235 @@
use serde_json::{json, Value};
use std::process::Command;
#[derive(Debug, Clone)]
pub struct IosDevice {
pub name: String,
pub udid: String,
pub state: String,
pub runtime: String,
pub is_real: bool,
}
pub fn list_simulators() -> Result<Vec<IosDevice>, String> {
let output = Command::new("xcrun")
.args(["simctl", "list", "devices", "--json"])
.output()
.map_err(|e| format!("Failed to run xcrun simctl: {}", e))?;
if !output.status.success() {
return Err("xcrun simctl failed. Xcode may not be installed.".to_string());
}
let json_str = String::from_utf8_lossy(&output.stdout);
let parsed: Value =
serde_json::from_str(&json_str).map_err(|e| format!("Failed to parse simctl: {}", e))?;
let mut devices = Vec::new();
if let Some(device_map) = parsed.get("devices").and_then(|v| v.as_object()) {
for (runtime, device_list) in device_map {
if let Some(arr) = device_list.as_array() {
for device in arr {
let name = device
.get("name")
.and_then(|v| v.as_str())
.unwrap_or("")
.to_string();
let udid = device
.get("udid")
.and_then(|v| v.as_str())
.unwrap_or("")
.to_string();
let state = device
.get("state")
.and_then(|v| v.as_str())
.unwrap_or("")
.to_string();
devices.push(IosDevice {
name,
udid,
state,
runtime: runtime.clone(),
is_real: false,
});
}
}
}
}
Ok(devices)
}
pub fn list_real_devices() -> Result<Vec<IosDevice>, String> {
let output = Command::new("xcrun")
.args(["xctrace", "list", "devices"])
.output()
.map_err(|e| format!("Failed to run xcrun xctrace: {}", e))?;
if !output.status.success() {
return Ok(Vec::new());
}
let stdout = String::from_utf8_lossy(&output.stdout);
let mut devices = Vec::new();
let mut in_devices = false;
for line in stdout.lines() {
let trimmed = line.trim();
if trimmed.starts_with("== Devices ==") {
in_devices = true;
continue;
}
if trimmed.starts_with("== Simulators ==") {
break;
}
if !in_devices || trimmed.is_empty() {
continue;
}
// Format: "Device Name (OS Version) (UDID)"
if let Some(udid_start) = trimmed.rfind('(') {
let udid_end = trimmed.len() - 1;
let udid = &trimmed[udid_start + 1..udid_end];
// Validate it looks like a UDID (contains hyphens)
if udid.contains('-') && udid.len() > 20 {
let name_part = trimmed[..udid_start].trim();
let name = if let Some(paren_pos) = name_part.rfind('(') {
name_part[..paren_pos].trim().to_string()
} else {
name_part.to_string()
};
devices.push(IosDevice {
name,
udid: udid.to_string(),
state: "Connected".to_string(),
runtime: String::new(),
is_real: true,
});
}
}
}
Ok(devices)
}
pub fn list_all_devices() -> Result<Vec<IosDevice>, String> {
let mut all = list_simulators().unwrap_or_default();
all.extend(list_real_devices().unwrap_or_default());
Ok(all)
}
pub fn boot_simulator(udid: &str) -> Result<(), String> {
let output = Command::new("xcrun")
.args(["simctl", "boot", udid])
.output()
.map_err(|e| format!("Failed to boot simulator: {}", e))?;
if !output.status.success() {
let stderr = String::from_utf8_lossy(&output.stderr);
if stderr.contains("current state: Booted") {
return Ok(());
}
return Err(format!("Failed to boot simulator {}: {}", udid, stderr));
}
Ok(())
}
pub fn shutdown_simulator(udid: &str) -> Result<(), String> {
let output = Command::new("xcrun")
.args(["simctl", "shutdown", udid])
.output()
.map_err(|e| format!("Failed to shutdown simulator: {}", e))?;
if !output.status.success() {
let stderr = String::from_utf8_lossy(&output.stderr);
if stderr.contains("current state: Shutdown") {
return Ok(());
}
return Err(format!("Failed to shutdown simulator {}: {}", udid, stderr));
}
Ok(())
}
pub fn select_device(device_name: Option<&str>, udid: Option<&str>) -> Result<IosDevice, String> {
if let Some(u) = udid {
let devices = list_all_devices()?;
return devices
.into_iter()
.find(|d| d.udid == u)
.ok_or_else(|| format!("Device with UDID '{}' not found", u));
}
if let Some(name) = device_name {
let devices = list_all_devices()?;
return devices
.into_iter()
.find(|d| d.name.to_lowercase().contains(&name.to_lowercase()))
.ok_or_else(|| format!("Device '{}' not found", name));
}
// Default: prefer most recent iPhone, prefer Pro
let devices = list_simulators()?;
let iphone_devices: Vec<&IosDevice> = devices
.iter()
.filter(|d| d.name.starts_with("iPhone"))
.collect();
if iphone_devices.is_empty() {
return devices
.into_iter()
.next()
.ok_or("No iOS simulators found".to_string());
}
// Prefer Pro models
if let Some(pro) = iphone_devices.iter().find(|d| d.name.contains("Pro")) {
return Ok((*pro).clone());
}
Ok((*iphone_devices.last().unwrap()).clone())
}
pub fn to_device_json(devices: &[IosDevice]) -> Value {
let list: Vec<Value> = devices
.iter()
.map(|d| {
json!({
"name": d.name,
"udid": d.udid,
"state": d.state,
"runtime": d.runtime,
"isReal": d.is_real,
})
})
.collect();
json!({ "devices": list })
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_ios_device_struct() {
let device = IosDevice {
name: "iPhone 15 Pro".to_string(),
udid: "ABC-123".to_string(),
state: "Booted".to_string(),
runtime: "iOS-17-0".to_string(),
is_real: false,
};
assert_eq!(device.name, "iPhone 15 Pro");
assert!(!device.is_real);
}
#[test]
fn test_to_device_json() {
let devices = vec![IosDevice {
name: "Test".to_string(),
udid: "123".to_string(),
state: "Shutdown".to_string(),
runtime: "iOS-17".to_string(),
is_real: false,
}];
let json = to_device_json(&devices);
assert!(json.get("devices").unwrap().as_array().unwrap().len() == 1);
}
}
+6
View File
@@ -0,0 +1,6 @@
pub mod appium;
pub mod backend;
pub mod client;
pub mod ios;
pub mod safari;
pub mod types;
+80
View File
@@ -0,0 +1,80 @@
use std::path::PathBuf;
use std::process::{Child, Command, Stdio};
use std::time::Duration;
pub struct SafariDriverProcess {
child: Child,
pub port: u16,
}
impl SafariDriverProcess {
pub fn kill(&mut self) {
let _ = self.child.kill();
let _ = self.child.wait();
}
}
impl Drop for SafariDriverProcess {
fn drop(&mut self) {
self.kill();
}
}
pub fn find_safaridriver() -> Option<PathBuf> {
let candidates = ["/usr/bin/safaridriver"];
for c in &candidates {
let p = PathBuf::from(c);
if p.exists() {
return Some(p);
}
}
// Try PATH
if let Ok(output) = Command::new("which").arg("safaridriver").output() {
if output.status.success() {
let path = String::from_utf8_lossy(&output.stdout).trim().to_string();
if !path.is_empty() {
return Some(PathBuf::from(path));
}
}
}
None
}
pub fn launch_safaridriver(port: u16) -> Result<SafariDriverProcess, String> {
let driver_path = find_safaridriver()
.ok_or("safaridriver not found. Safari WebDriver requires macOS with Safari.")?;
let child = Command::new(&driver_path)
.arg("--port")
.arg(port.to_string())
.stdin(Stdio::null())
.stdout(Stdio::null())
.stderr(Stdio::null())
.spawn()
.map_err(|e| format!("Failed to launch safaridriver: {}", e))?;
// Wait for driver to be ready
std::thread::sleep(Duration::from_millis(500));
Ok(SafariDriverProcess { child, port })
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_find_safaridriver() {
// Only check on macOS
if cfg!(target_os = "macos") {
let result = find_safaridriver();
// Don't assert Some since it may not be enabled
if let Some(path) = result {
assert!(path.exists());
}
}
}
}
+97
View File
@@ -0,0 +1,97 @@
use serde::{Deserialize, Serialize};
use serde_json::Value;
#[derive(Debug, Serialize)]
#[serde(rename_all = "camelCase")]
pub struct NewSessionRequest {
pub capabilities: Capabilities,
}
#[derive(Debug, Serialize)]
#[serde(rename_all = "camelCase")]
pub struct Capabilities {
pub always_match: Value,
}
#[derive(Debug, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct SessionResponse {
pub value: SessionValue,
}
#[derive(Debug, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct SessionValue {
pub session_id: String,
pub capabilities: Value,
}
#[derive(Debug, Deserialize)]
pub struct WebDriverResponse {
pub value: Value,
}
#[derive(Debug, Deserialize)]
pub struct WebDriverError {
pub error: String,
pub message: String,
}
#[derive(Debug, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct ElementResponse {
pub value: ElementValue,
}
#[derive(Debug, Deserialize)]
pub struct ElementValue {
#[serde(rename = "element-6066-11e4-a52e-4f735466cecf")]
pub element_id: Option<String>,
#[serde(rename = "ELEMENT")]
pub element_legacy: Option<String>,
}
impl ElementValue {
pub fn id(&self) -> Option<&str> {
self.element_id
.as_deref()
.or(self.element_legacy.as_deref())
}
}
#[derive(Debug, Serialize)]
pub struct FindElementRequest {
pub using: String,
pub value: String,
}
#[derive(Debug, Serialize)]
pub struct ExecuteScriptRequest {
pub script: String,
pub args: Vec<Value>,
}
#[derive(Debug, Serialize)]
#[serde(rename_all = "camelCase")]
pub struct CookieRequest {
pub cookie: CookieData,
}
#[derive(Debug, Serialize, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct CookieData {
pub name: String,
pub value: String,
#[serde(skip_serializing_if = "Option::is_none")]
pub domain: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub path: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub secure: Option<bool>,
#[serde(skip_serializing_if = "Option::is_none")]
pub http_only: Option<bool>,
#[serde(skip_serializing_if = "Option::is_none")]
pub expiry: Option<u64>,
#[serde(skip_serializing_if = "Option::is_none")]
pub same_site: Option<String>,
}
+2573 -39
View File
File diff suppressed because it is too large Load Diff
+49
View File
@@ -0,0 +1,49 @@
use std::sync::{Mutex, MutexGuard};
/// Global mutex shared across all test modules to prevent parallel tests from
/// interfering with each other when mutating environment variables.
pub static ENV_MUTEX: Mutex<()> = Mutex::new(());
/// RAII guard that locks [`ENV_MUTEX`] and restores environment variables on drop.
pub struct EnvGuard<'a> {
_lock: MutexGuard<'a, ()>,
vars: Vec<(String, Option<String>)>,
}
impl<'a> EnvGuard<'a> {
pub fn new(var_names: &[&str]) -> Self {
let lock = ENV_MUTEX.lock().unwrap();
let vars = var_names
.iter()
.map(|&name| (name.to_string(), std::env::var(name).ok()))
.collect();
Self { _lock: lock, vars }
}
pub fn set(&self, name: &str, value: &str) {
debug_assert!(
self.vars.iter().any(|(n, _)| n == name),
"EnvGuard::set called with unregistered var: {name}"
);
std::env::set_var(name, value);
}
pub fn remove(&self, name: &str) {
debug_assert!(
self.vars.iter().any(|(n, _)| n == name),
"EnvGuard::remove called with unregistered var: {name}"
);
std::env::remove_var(name);
}
}
impl Drop for EnvGuard<'_> {
fn drop(&mut self) {
for (name, value) in &self.vars {
match value {
Some(v) => std::env::set_var(name, v),
None => std::env::remove_var(name),
}
}
}
}
+15
View File
@@ -0,0 +1,15 @@
/// Check if a session name is valid (alphanumeric, hyphens, and underscores only)
pub fn is_valid_session_name(name: &str) -> bool {
!name.is_empty()
&& name
.chars()
.all(|c| c.is_alphanumeric() || c == '-' || c == '_')
}
/// Generate error message for invalid session name
pub fn session_name_error(name: &str) -> String {
format!(
"Invalid session name '{}'. Only alphanumeric characters, hyphens, and underscores are allowed.",
name
)
}
+1
View File
@@ -0,0 +1 @@
{"v":1}
+6
View File
@@ -0,0 +1,6 @@
{
"git": {
"sha1": "7cb6c7d950c040b2198da553140e1b5e8b6ac682"
},
"path_in_vcs": "crates/zune-jpeg"
}
+1
View File
@@ -0,0 +1 @@
/target
+79
View File
@@ -0,0 +1,79 @@
# Benchmarks of popular jpeg libraries
Here I compare how long it takes popular JPEG decoders to decode the below 7680*4320 image
of (now defunct ?) [Cutefish OS](https://en.cutefishos.com/) default wallpaper.
![img](benches/images/speed_bench.jpg)
## About benchmarks
Benchmarks are weird, especially IO & multi-threaded programs. This library uses both of the above hence performance may
vary.
For best results shut down your machine, go take coffee, think about life and how it came to be and why people should
save the environment.
Then power up your machine, if it's a laptop connect it to a power supply and if there is a setting for performance
mode, tweak it.
Then run.
## Benchmarks vs real world usage
Real world usage may vary.
Notice that I'm using a large image but probably most decoding will be small to medium images.
To make the library thread safe, we do about 1.5-1.7x more allocations than libjpeg-turbo. Although, do note that the
allocations do not occur at ago, we allocate when needed and deallocate when not needed.
Do note if memory bandwidth is a limitation. This is not for you.
## Reproducibility
The benchmarks are carried out on my local machine with an AMD Ryzen 5 4500u
The benchmarks are reproducible.
To reproduce them
1. Clone this repository
2. Install rust(if you don't have it yet)
3. `cd` into the directory.
4. Run `cargo bench`
## Performance features of the three libraries
| feature | image-rs/jpeg-decoder | libjpeg-turbo | zune-jpeg |
|------------------------------|-----------------------|---------------|-----------|
| multithreaded | ✅ | ❌ | ❌ |
| platform specific intrinsics | ✅ | ✅ | ✅ |
- Image-rs/jpeg-decoder uses [rayon] under the hood but it's under a feature
flag.
- libjpeg-turbo uses hand-written asm for platform specific intrinsics, ported to
the most common architectures out there but falls back to scalar
code if it can't run in a platform.
# Finally benchmarks
[here]
## Notes
Benchmarks are ran at least once a week to catch regressions early and
are uploaded to Github pages.
Machine specs can be found on the other [landing page]
Benchmarks may not reflect real world usage(threads, other I/O machine bottlenecks)
[landing page]:https://etemesi254.github.io/posts/Zune-Benchmarks/
[here]:https://etemesi254.github.io/assets/criterion/report/index.html
[libjpeg-turbo]:https://github.com/libjpeg-turbo/libjpeg-turbo
[jpeg-decoder]:https://github.com/image-rs/jpeg-decoder
[rayon]:https://github.com/rayon-rs/rayon
+25
View File
@@ -0,0 +1,25 @@
# This file is automatically @generated by Cargo.
# It is not intended for manual editing.
version = 4
[[package]]
name = "log"
version = "0.4.29"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "5e5032e24019045c762d3c0f28f5b6b8bbf38563a65908389bf7978758920897"
[[package]]
name = "zune-core"
version = "0.5.1"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "cb8a0807f7c01457d0379ba880ba6322660448ddebc890ce29bb64da71fb40f9"
dependencies = [
"log",
]
[[package]]
name = "zune-jpeg"
version = "0.5.12"
dependencies = [
"zune-core",
]
+67
View File
@@ -0,0 +1,67 @@
# THIS FILE IS AUTOMATICALLY GENERATED BY CARGO
#
# When uploading crates to the registry Cargo will automatically
# "normalize" Cargo.toml files for maximal compatibility
# with all versions of Cargo and also rewrite `path` dependencies
# to registry (e.g., crates.io) dependencies.
#
# If you are reading this file be aware that the original Cargo.toml
# will likely look very different (and much more reasonable).
# See Cargo.toml.orig for the original contents.
[package]
edition = "2021"
rust-version = "1.75.0"
name = "zune-jpeg"
version = "0.5.12"
authors = ["caleb <etemesicaleb@gmail.com>"]
build = false
exclude = [
"/benches/images/*",
"/tests/*",
"/.idea/*",
"/.gradle/*",
"/test-images/*",
"fuzz/*",
]
autolib = false
autobins = false
autoexamples = false
autotests = false
autobenches = false
description = "A fast, correct and safe jpeg decoder"
readme = "README.md"
keywords = [
"jpeg",
"jpeg-decoder",
"decoder",
]
categories = ["multimedia::images"]
license = "MIT OR Apache-2.0 OR Zlib"
repository = "https://github.com/etemesi254/zune-image/tree/dev/crates/zune-jpeg"
[features]
default = [
"x86",
"neon",
"std",
]
log = ["zune-core/log"]
neon = []
portable_simd = []
std = ["zune-core/std"]
x86 = []
[lib]
name = "zune_jpeg"
path = "src/lib.rs"
[dependencies.zune-core]
version = "0.5.1"
[dev-dependencies]
[lints.rust.unexpected_cfgs]
level = "warn"
priority = 0
check-cfg = ["cfg(fuzzing)"]
+36
View File
@@ -0,0 +1,36 @@
[package]
name = "zune-jpeg"
version = "0.5.12"
rust-version = "1.75.0"
authors = ["caleb <etemesicaleb@gmail.com>"]
edition = "2021"
repository = "https://github.com/etemesi254/zune-image/tree/dev/crates/zune-jpeg"
license = "MIT OR Apache-2.0 OR Zlib"
keywords = ["jpeg", "jpeg-decoder", "decoder"]
categories = ["multimedia::images"]
exclude = ["/benches/images/*", "/tests/*", "/.idea/*", "/.gradle/*", "/test-images/*", "fuzz/*"]
description = "A fast, correct and safe jpeg decoder"
[lints.rust]
# Disable feature checker for fuzzing since it's used and cargo doesn't
# seem to recognise fuzzing
unexpected_cfgs = { level = "warn", check-cfg = ['cfg(fuzzing)'] }
# See more keys and their definitions at https://doc.rust-lang.org/cargo/reference/manifest.html
[features]
x86 = []
neon = []
std = ["zune-core/std"]
# NOTE: portable_simd requires Rust 1.87+
portable_simd = []
log = ["zune-core/log"]
default = ["x86", "neon", "std"]
[dependencies]
zune-core = { path = "../zune-core", version = "0.5.1" }
[dev-dependencies]
zune-ppm = { path = "../zune-ppm" }
+95
View File
@@ -0,0 +1,95 @@
## Version 0.5.7
- Move scalar idct to wrapping maths.
- Simd upsampling (mhils)
- Faster zero idct check (mhils)
## Version 0.5.6
- Better support for truncated images (by https://github.com/mhils)
- fix 4:1:0 chroma subsampling (by https://github.com/mhils)
- Fix some crashes
- Fix some bug on last pixel sampling
## Version 0.5.5
- Support direct conversion of Luma to RGBA
## Version 0.5.4
- Fix overriding color space when decoding Luma colorspace
## Version 0.5.3
- Fix some decoding of some images with markers in progressive segments, see https://github.com/etemesi254/zune-image/issues/295
## Version 0.5.1
- Fix decoding of particular images with a non-standard subsample, (
see https://github.com/etemesi254/zune-image/issues/291)
- Add better RGB color detection of images to match libjpeg and stb_image formats
-----
## Version 0.3.17
- Fix no-std compilation
## Version 0.3.16
- Add support for decoding to BGR and BGRA
## Version 0.3.14
- Add ability to parse exif and ICC chunk.
- Fix images with one component that were down-sampled.
### Version 0.3.13
- Allow decoding into pre-allocated buffer
- Clarify documentation
### Version 0.3.11
- Add guards for SSE and AVX code paths(allows compiling for platforms that do not support it)
### Version 0.3.0
- Overhaul to the whole decoder.
- Single threaded version
- Lightweight.
### Version 0.2.0
- New `ZuneJpegOptions` struct, this is the now recommended way to set up decoding options for
decoding
- Deprecated previous options setting functions.
- More code cleanups
- Fixed new bugs discovered by fuzzing
- Removed dependency on `num_cpu`
### Version 0.1.5
- Allow user to set memory limits in during decoding explicitly via `set_limits`
- Fixed some bugs discovered by fuzzing
- Correctly handle small images less than 16 pixels
- Gracefully handle incorrectly sampled images.
### Version 0.1.4
- Remove all `unsafe` instances except platform dependent intrinsics.
- Numerous bug fixes identified by fuzzing.
- Expose `ImageInfo` to the crate root.
### Version 0.1.3
- Fix numerous panics found by fuzzing(thanks to @[Shnatsel] for the corpus)
- Add new method `set_num_threads` that allows one to explicitly set the number of threads to use to decode the image.
### Version 0.1.2
- Add more sub checks, contributed by @[5225225]
- Privatize some modules.
### Version 0.1.1
- Fix rgba/rgbx decoding when avx optimized functions were used
- Initial support for fuzzing
- Remove `align_alloc` method which was unsound (Thanks to @[HeroicKatora] for pointing that out)
[Shnatsel]:https://github.com/Shnatsel
[HeroicKatora]:https://github.com/HeroicKatora
[5225225]:https://github.com/5225225
+201
View File
@@ -0,0 +1,201 @@
Apache License
Version 2.0, January 2004
http://www.apache.org/licenses/
TERMS AND CONDITIONS FOR USE, REPRODUCTION, AND DISTRIBUTION
1. Definitions.
"License" shall mean the terms and conditions for use, reproduction,
and distribution as defined by Sections 1 through 9 of this document.
"Licensor" shall mean the copyright owner or entity authorized by
the copyright owner that is granting the License.
"Legal Entity" shall mean the union of the acting entity and all
other entities that control, are controlled by, or are under common
control with that entity. For the purposes of this definition,
"control" means (i) the power, direct or indirect, to cause the
direction or management of such entity, whether by contract or
otherwise, or (ii) ownership of fifty percent (50%) or more of the
outstanding shares, or (iii) beneficial ownership of such entity.
"You" (or "Your") shall mean an individual or Legal Entity
exercising permissions granted by this License.
"Source" form shall mean the preferred form for making modifications,
including but not limited to software source code, documentation
source, and configuration files.
"Object" form shall mean any form resulting from mechanical
transformation or translation of a Source form, including but
not limited to compiled object code, generated documentation,
and conversions to other media types.
"Work" shall mean the work of authorship, whether in Source or
Object form, made available under the License, as indicated by a
copyright notice that is included in or attached to the work
(an example is provided in the Appendix below).
"Derivative Works" shall mean any work, whether in Source or Object
form, that is based on (or derived from) the Work and for which the
editorial revisions, annotations, elaborations, or other modifications
represent, as a whole, an original work of authorship. For the purposes
of this License, Derivative Works shall not include works that remain
separable from, or merely link (or bind by name) to the interfaces of,
the Work and Derivative Works thereof.
"Contribution" shall mean any work of authorship, including
the original version of the Work and any modifications or additions
to that Work or Derivative Works thereof, that is intentionally
submitted to Licensor for inclusion in the Work by the copyright owner
or by an individual or Legal Entity authorized to submit on behalf of
the copyright owner. For the purposes of this definition, "submitted"
means any form of electronic, verbal, or written communication sent
to the Licensor or its representatives, including but not limited to
communication on electronic mailing lists, source code control systems,
and issue tracking systems that are managed by, or on behalf of, the
Licensor for the purpose of discussing and improving the Work, but
excluding communication that is conspicuously marked or otherwise
designated in writing by the copyright owner as "Not a Contribution."
"Contributor" shall mean Licensor and any individual or Legal Entity
on behalf of whom a Contribution has been received by Licensor and
subsequently incorporated within the Work.
2. Grant of Copyright License. Subject to the terms and conditions of
this License, each Contributor hereby grants to You a perpetual,
worldwide, non-exclusive, no-charge, royalty-free, irrevocable
copyright license to reproduce, prepare Derivative Works of,
publicly display, publicly perform, sublicense, and distribute the
Work and such Derivative Works in Source or Object form.
3. Grant of Patent License. Subject to the terms and conditions of
this License, each Contributor hereby grants to You a perpetual,
worldwide, non-exclusive, no-charge, royalty-free, irrevocable
(except as stated in this section) patent license to make, have made,
use, offer to sell, sell, import, and otherwise transfer the Work,
where such license applies only to those patent claims licensable
by such Contributor that are necessarily infringed by their
Contribution(s) alone or by combination of their Contribution(s)
with the Work to which such Contribution(s) was submitted. If You
institute patent litigation against any entity (including a
cross-claim or counterclaim in a lawsuit) alleging that the Work
or a Contribution incorporated within the Work constitutes direct
or contributory patent infringement, then any patent licenses
granted to You under this License for that Work shall terminate
as of the date such litigation is filed.
4. Redistribution. You may reproduce and distribute copies of the
Work or Derivative Works thereof in any medium, with or without
modifications, and in Source or Object form, provided that You
meet the following conditions:
(a) You must give any other recipients of the Work or
Derivative Works a copy of this License; and
(b) You must cause any modified files to carry prominent notices
stating that You changed the files; and
(c) You must retain, in the Source form of any Derivative Works
that You distribute, all copyright, patent, trademark, and
attribution notices from the Source form of the Work,
excluding those notices that do not pertain to any part of
the Derivative Works; and
(d) If the Work includes a "NOTICE" text file as part of its
distribution, then any Derivative Works that You distribute must
include a readable copy of the attribution notices contained
within such NOTICE file, excluding those notices that do not
pertain to any part of the Derivative Works, in at least one
of the following places: within a NOTICE text file distributed
as part of the Derivative Works; within the Source form or
documentation, if provided along with the Derivative Works; or,
within a display generated by the Derivative Works, if and
wherever such third-party notices normally appear. The contents
of the NOTICE file are for informational purposes only and
do not modify the License. You may add Your own attribution
notices within Derivative Works that You distribute, alongside
or as an addendum to the NOTICE text from the Work, provided
that such additional attribution notices cannot be construed
as modifying the License.
You may add Your own copyright statement to Your modifications and
may provide additional or different license terms and conditions
for use, reproduction, or distribution of Your modifications, or
for any such Derivative Works as a whole, provided Your use,
reproduction, and distribution of the Work otherwise complies with
the conditions stated in this License.
5. Submission of Contributions. Unless You explicitly state otherwise,
any Contribution intentionally submitted for inclusion in the Work
by You to the Licensor shall be under the terms and conditions of
this License, without any additional terms or conditions.
Notwithstanding the above, nothing herein shall supersede or modify
the terms of any separate license agreement you may have executed
with Licensor regarding such Contributions.
6. Trademarks. This License does not grant permission to use the trade
names, trademarks, service marks, or product names of the Licensor,
except as required for reasonable and customary use in describing the
origin of the Work and reproducing the content of the NOTICE file.
7. Disclaimer of Warranty. Unless required by applicable law or
agreed to in writing, Licensor provides the Work (and each
Contributor provides its Contributions) on an "AS IS" BASIS,
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or
implied, including, without limitation, any warranties or conditions
of TITLE, NON-INFRINGEMENT, MERCHANTABILITY, or FITNESS FOR A
PARTICULAR PURPOSE. You are solely responsible for determining the
appropriateness of using or redistributing the Work and assume any
risks associated with Your exercise of permissions under this License.
8. Limitation of Liability. In no event and under no legal theory,
whether in tort (including negligence), contract, or otherwise,
unless required by applicable law (such as deliberate and grossly
negligent acts) or agreed to in writing, shall any Contributor be
liable to You for damages, including any direct, indirect, special,
incidental, or consequential damages of any character arising as a
result of this License or out of the use or inability to use the
Work (including but not limited to damages for loss of goodwill,
work stoppage, computer failure or malfunction, or any and all
other commercial damages or losses), even if such Contributor
has been advised of the possibility of such damages.
9. Accepting Warranty or Additional Liability. While redistributing
the Work or Derivative Works thereof, You may choose to offer,
and charge a fee for, acceptance of support, warranty, indemnity,
or other liability obligations and/or rights consistent with this
License. However, in accepting such obligations, You may act only
on Your own behalf and on Your sole responsibility, not on behalf
of any other Contributor, and only if You agree to indemnify,
defend, and hold each Contributor harmless for any liability
incurred by, or claims asserted against, such Contributor by reason
of your accepting any such warranty or additional liability.
END OF TERMS AND CONDITIONS
APPENDIX: How to apply the Apache License to your work.
To apply the Apache License to your work, attach the following
boilerplate notice, with the fields enclosed by brackets "[]"
replaced with your own identifying information. (Don't include
the brackets!) The text should be enclosed in the appropriate
comment syntax for the file format. We also recommend that a
file or class name and description of purpose be included on the
same "printed page" as the copyright notice for easier
identification within third-party archives.
Copyright [yyyy] [name of copyright owner]
Licensed under the Apache License, Version 2.0 (the "License");
you may not use this file except in compliance with the License.
You may obtain a copy of the License at
http://www.apache.org/licenses/LICENSE-2.0
Unless required by applicable law or agreed to in writing, software
distributed under the License is distributed on an "AS IS" BASIS,
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
See the License for the specific language governing permissions and
limitations under the License.
+21
View File
@@ -0,0 +1,21 @@
MIT License
Copyright (c) zune-image developers
Permission is hereby granted, free of charge, to any person obtaining a copy
of this software and associated documentation files (the "Software"), to deal
in the Software without restriction, including without limitation the rights
to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
copies of the Software, and to permit persons to whom the Software is
furnished to do so, subject to the following conditions:
The above copyright notice and this permission notice shall be included in all
copies or substantial portions of the Software.
THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
SOFTWARE.
+19
View File
@@ -0,0 +1,19 @@
zlib License
(C) zune-image developers
This software is provided 'as-is', without any express or implied
warranty. In no event will the authors be held liable for any damages
arising from the use of this software.
Permission is granted to anyone to use this software for any purpose,
including commercial applications, and to alter it and redistribute it
freely, subject to the following restrictions:
1. The origin of this software must not be misrepresented; you must not
claim that you wrote the original software. If you use this software
in a product, an acknowledgment in the product documentation would be
appreciated but is not required.
2. Altered source versions must be plainly marked as such, and must not be
misrepresented as being the original software.
3. This notice may not be removed or altered from any source distribution.
+104
View File
@@ -0,0 +1,104 @@
# Zune-JPEG
A fast, correct and safe jpeg decoder in pure Rust.
## Usage
The library provides a simple-to-use API for jpeg decoding
and an ability to add options to influence decoding.
### Example
```Rust
// Import the library
use zune_jpeg::JpegDecoder;
use std::fs::read;
fn main()->Result<(),DecoderErrors> {
// load some jpeg data
let data = read("cat.jpg").unwrap();
// create a decoder
let mut decoder = JpegDecoder::new(&data);
// decode the file
let pixels = decoder.decode()?;
}
```
The decoder supports more manipulations via `DecoderOptions`,
see additional documentation in the library.
## Goals
The implementation aims to have the following goals achieved,
in order of importance
1. Safety - Do not segfault on errors or invalid input. Panics are okay, but
should be fixed when reported. `unsafe` is only used for SIMD intrinsics,
and can be turned off entirely both at compile time and at runtime.
2. Speed - Get the data as quickly as possible, which means
1. Platform intrinsics code where justifiable
2. Carefully written platform independent code that allows the
compiler to vectorize it.
3. Regression tests.
4. Watch the memory usage of the program
3. Usability - Provide utility functions like different color conversions functions.
## Non-Goals
- Bit identical results with libjpeg/libjpeg-turbo will never be an aim of this library.
Jpeg is a lossy format with very few parts specified by the standard
(i.e it doesn't give a reference upsampling and color conversion algorithm)
## Features
- [x] A Pretty fast 8*8 integer IDCT.
- [x] Fast Huffman Decoding
- [x] Fast color convert functions.
- [x] Support for extended colorspaces like GrayScale and RGBA
- [X] Single-threaded decoding.
- [X] Support for four component JPEGs, and esoteric color schemes like CYMK
- [X] Support for `no_std`
- [X] BGR/BGRA decoding support.
## Crate Features
| feature | on | Capabilities |
|---------|-----|---------------------------------------------------------------------------------------------|
| `x86` | yes | Enables `x86` specific instructions, specifically `avx` and `sse` for accelerated decoding. |
| `std` | yes | Enable linking to the `std` crate |
Note that the `x86` features are automatically disabled on platforms that aren't x86 during compile
time hence there is no need to disable them explicitly if you are targeting such a platform.
## Using in a `no_std` environment
The crate can be used in a `no_std` environment with the `alloc` feature.
But one is required to link to a working allocator for whatever environment the decoder
will be running on
## Debug vs release
The decoder heavily relies on platform specific intrinsics, namely AVX2 and SSE to gain speed-ups in decoding,
but they [perform poorly](https://godbolt.org/z/vPq57z13b) in debug builds. To get reasonable performance even
when compiling your program in debug mode, add this to your `Cargo.toml`:
```toml
# `zune-jpeg` package will be always built with optimizations
[profile.dev.package.zune-jpeg]
opt-level = 3
```
## Benchmarks
The library tries to be at fast as [libjpeg-turbo] while being as safe as possible.
Platform specific intrinsics help get speed up intensive operations ensuring we can almost
match [libjpeg-turbo] speeds but speeds are always +- 10 ms of this library.
For more up-to-date benchmarks, see the online repo with
benchmarks [here](https://etemesi254.github.io/assets/criterion/report/index.html)
[libjpeg-turbo]:https://github.com/libjpeg-turbo/libjpeg-turbo/
[image-rs/jpeg-decoder]:https://github.com/image-rs/jpeg-decoder/tree/master/src
+811
View File
@@ -0,0 +1,811 @@
/*
* Copyright (c) 2023.
*
* This software is free software;
*
* You can redistribute it or modify it under terms of the MIT, Apache License or Zlib license
*/
#![allow(
clippy::if_not_else,
clippy::similar_names,
clippy::inline_always,
clippy::doc_markdown,
clippy::cast_sign_loss,
clippy::cast_possible_truncation
)]
//! This file exposes a single struct that can decode a huffman encoded
//! Bitstream in a JPEG file
//!
//! This code is optimized for speed.
//! It's meant to be super duper super fast, because everyone else depends on this being fast.
//! It's (annoyingly) serial hence we cant use parallel bitstreams(it's variable length coding.)
//!
//! Furthermore, on the case of refills, we have to do bytewise processing because the standard decided
//! that we want to support markers in the middle of streams(seriously few people use RST markers).
//!
//! So we pull in all optimization steps:
//! - use `inline[always]`? ✅ ,
//! - pre-execute most common cases ✅,
//! - add random comments ✅
//! - fast paths ✅.
//!
//! Speed-wise: It is probably the fastest JPEG BitStream decoder to ever sail the seven seas because of
//! a couple of optimization tricks.
//! 1. Fast refills from libjpeg-turbo
//! 2. As few as possible branches in decoder fast paths.
//! 3. Accelerated AC table decoding borrowed from stb_image.h written by Fabian Gissen (@ rygorous),
//! improved by me to handle more cases.
//! 4. Safe and extensible routines(e.g. cool ways to eliminate bounds check)
//! 5. No unsafe here
//!
//! Readability comes as a second priority(I tried with variable names this time, and we are wayy better than libjpeg).
//!
//! Anyway if you are reading this it means your cool and I hope you get whatever part of the code you are looking for
//! (or learn something cool)
//!
//! Knock yourself out.
use alloc::format;
use alloc::string::ToString;
use core::cmp::min;
use zune_core::bytestream::{ZByteReaderTrait, ZReader};
use crate::errors::DecodeErrors;
use crate::huffman::{HuffmanTable, HUFF_LOOKAHEAD};
use crate::marker::Marker;
use crate::mcu::DCT_BLOCK;
use crate::misc::UN_ZIGZAG;
macro_rules! decode_huff {
($stream:tt,$symbol:tt,$table:tt) => {
let mut code_length = $symbol >> HUFF_LOOKAHEAD;
($symbol) &= (1 << HUFF_LOOKAHEAD) - 1;
if code_length > i32::from(HUFF_LOOKAHEAD)
{
// if the symbol cannot be resolved in the first HUFF_LOOKAHEAD bits,
// we know it lies somewhere between HUFF_LOOKAHEAD and 16 bits since jpeg imposes 16 bit
// limit, we can therefore look 16 bits ahead and try to resolve the symbol
// starting from 1+HUFF_LOOKAHEAD bits.
$symbol = ($stream).peek_bits::<16>() as i32;
// (Credits to Sean T. Barrett stb library for this optimization)
// maxcode is pre-shifted 16 bytes long so that it has (16-code_length)
// zeroes at the end hence we do not need to shift in the inner loop.
while code_length < 17{
if $symbol < $table.maxcode[code_length as usize] {
break;
}
code_length += 1;
}
if code_length == 17{
// symbol could not be decoded.
//
// We may think, lets fake zeroes, noo
// panic, because Huffman codes are sensitive, probably everything
// after this will be corrupt, so no need to continue.
// panic!("Bad Huffman code length");
return Err(DecodeErrors::Format(format!("Bad Huffman Code 0x{:X}, corrupt JPEG",$symbol)))
}
$symbol >>= (16-code_length);
($symbol) = i32::from(
($table).values
[(($symbol + ($table).offset[code_length as usize]) & 0xFF) as usize],
);
}
if code_length> i32::from(($stream).bits_left){
return Err(DecodeErrors::Format(format!("Code length {code_length} more than bits left {}",($stream).bits_left)))
}
// drop bits read
($stream).drop_bits(code_length as u8);
};
}
/// A `BitStream` struct, a bit by bit reader with super powers
///
#[rustfmt::skip]
pub(crate) struct BitStream {
/// A MSB type buffer that is used for some certain operations
pub buffer: u64,
/// A TOP aligned MSB type buffer that is used to accelerate some operations like
/// peek_bits and get_bits.
///
/// By top aligned, I mean the top bit (63) represents the top bit in the buffer.
aligned_buffer: u64,
/// Tell us the bits left the two buffer
pub(crate) bits_left: u8,
/// Did we find a marker(RST/EOF) during decoding?
pub marker: Option<Marker>,
/// An i16 with the bit corresponding to successive_low set to 1, others 0.
pub successive_low_mask: i16,
spec_start: u8,
spec_end: u8,
pub eob_run: i32,
pub overread_by: usize,
/// True if we have seen end of image marker.
/// Don't read anything after that.
pub seen_eoi: bool,
}
impl BitStream {
/// Create a new BitStream
#[rustfmt::skip]
pub(crate) const fn new() -> BitStream {
BitStream {
buffer: 0,
aligned_buffer: 0,
bits_left: 0,
marker: None,
successive_low_mask: 1,
spec_start: 0,
spec_end: 0,
eob_run: 0,
overread_by: 0,
seen_eoi: false,
}
}
/// Create a new Bitstream for progressive decoding
#[allow(clippy::redundant_field_names)]
#[rustfmt::skip]
pub(crate) fn new_progressive(al: u8, spec_start: u8, spec_end: u8) -> BitStream {
BitStream {
buffer: 0,
aligned_buffer: 0,
bits_left: 0,
marker: None,
successive_low_mask: 1i16 << al,
spec_start: spec_start,
spec_end: spec_end,
eob_run: 0,
overread_by: 0,
seen_eoi: false,
}
}
/// Refill the bit buffer by (a maximum of) 32 bits
///
/// # Arguments
/// - `reader`:`&mut BufReader<R>`: A mutable reference to an underlying
/// File/Memory buffer containing a valid JPEG stream
///
/// This function will only refill if `self.count` is less than 32
#[inline(always)] // to many call sites? ( perf improvement by 4%)
pub fn refill<T>(&mut self, reader: &mut ZReader<T>) -> Result<bool, DecodeErrors>
where
T: ZByteReaderTrait
{
/// Macro version of a single byte refill.
/// Arguments
/// buffer-> our io buffer, because rust macros cannot get values from
/// the surrounding environment bits_left-> number of bits left
/// to full refill
macro_rules! refill {
($buffer:expr,$byte:expr,$bits_left:expr) => {
// read a byte from the stream
$byte = u64::from(reader.read_u8());
self.overread_by += usize::from(reader.eof()?);
// append to the buffer
// JPEG is a MSB type buffer so that means we append this
// to the lower end (0..8) of the buffer and push the rest bits above..
$buffer = ($buffer << 8) | $byte;
// Increment bits left
$bits_left += 8;
// Check for special case of OxFF, to see if it's a stream or a marker
if $byte == 0xff {
// read next byte
let mut next_byte = u64::from(reader.read_u8());
// Byte snuffing, if we encounter byte snuff, we skip the byte
if next_byte != 0x00 {
// skip that byte we read
while next_byte == 0xFF {
next_byte = u64::from(reader.read_u8());
}
if next_byte != 0x00 {
// Undo the byte append and return
$buffer >>= 8;
$bits_left -= 8;
if $bits_left != 0 {
self.aligned_buffer = $buffer << (64 - $bits_left);
}
let marker = Marker::from_u8(next_byte as u8);
self.marker = marker;
if let Some(Marker::UNKNOWN(_)) = marker{
return Err(DecodeErrors::Format("Unknown marker in bit stream".to_string()));
}
if next_byte == 0xD9 {
// special handling for eoi, fill some bytes,even if its zero,
// removes some panics
self.buffer <<= 8;
self.bits_left += 8;
self.aligned_buffer = self.buffer << (64 - self.bits_left);
}
return Ok(false);
}
}
}
};
}
// 32 bits is enough for a decode(16 bits) and receive_extend(max 16 bits)
if self.bits_left < 32 {
if self.marker.is_some() || self.overread_by > 0 || self.seen_eoi {
// found a marker, or we are in EOI
// also we are in over-reading mode, where we fill it with zeroes
// fill with zeroes
self.buffer <<= 32;
self.bits_left += 32;
self.aligned_buffer = self.buffer << (64 - self.bits_left);
return Ok(true);
}
// we optimize for the case where we don't have 255 in the stream and have 4 bytes left
// as it is the common case
//
// so we always read 4 bytes, if read_fixed_bytes errors out, the cursor is
// guaranteed not to advance in case of failure (is this true), so
// we revert the read later on (if we have 255), if this fails, we use the normal
// byte at a time read
if let Ok(bytes) = reader.read_fixed_bytes_or_error::<4>() {
// we have 4 bytes to spare, read the 4 bytes into a temporary buffer
// create buffer
let msb_buf = u32::from_be_bytes(bytes);
// check if we have 0xff
if !has_byte(msb_buf, 255) {
self.bits_left += 32;
self.buffer <<= 32;
self.buffer |= u64::from(msb_buf);
self.aligned_buffer = self.buffer << (64 - self.bits_left);
return Ok(true);
}
reader.rewind(4)?;
}
// This serves two reasons,
// 1: Make clippy shut up
// 2: Favour register reuse
let mut byte;
// 4 refills, if all succeed the stream should contain enough bits to decode a
// value
refill!(self.buffer, byte, self.bits_left);
refill!(self.buffer, byte, self.bits_left);
refill!(self.buffer, byte, self.bits_left);
refill!(self.buffer, byte, self.bits_left);
// Construct an MSB buffer whose top bits are the bitstream we are currently holding.
self.aligned_buffer = self.buffer << (64 - self.bits_left);
}
return Ok(true);
}
/// Decode the DC coefficient in a MCU block.
///
/// The decoded coefficient is written to `dc_prediction`
///
#[allow(
clippy::cast_possible_truncation,
clippy::cast_sign_loss,
clippy::unwrap_used
)]
#[inline(always)]
fn decode_dc<T>(
&mut self, reader: &mut ZReader<T>, dc_table: &HuffmanTable, dc_prediction: &mut i32
) -> Result<bool, DecodeErrors>
where
T: ZByteReaderTrait
{
let (mut symbol, r);
if self.bits_left < 32 {
self.refill(reader)?;
};
// look a head HUFF_LOOKAHEAD bits into the bitstream
symbol = self.peek_bits::<HUFF_LOOKAHEAD>();
symbol = dc_table.lookup[symbol as usize];
decode_huff!(self, symbol, dc_table);
if symbol != 0 {
r = self.get_bits(symbol as u8);
symbol = huff_extend(r, symbol);
}
// Update DC prediction
*dc_prediction = dc_prediction.wrapping_add(symbol);
return Ok(true);
}
/// Like `decode_dc` but we do not need the result of the component, we only want to remove it
/// from the bitstream of the MCU.
fn discard_dc<T>(
&mut self, reader: &mut ZReader<T>, dc_table: &HuffmanTable
) -> Result<bool, DecodeErrors>
where
T: ZByteReaderTrait
{
let mut symbol;
if self.bits_left < 32 {
self.refill(reader)?;
};
// look a head HUFF_LOOKAHEAD bits into the bitstream
symbol = self.peek_bits::<HUFF_LOOKAHEAD>();
symbol = dc_table.lookup[symbol as usize];
decode_huff!(self, symbol, dc_table);
if symbol != 0 {
let _ = self.get_bits(symbol as u8);
}
return Ok(true);
}
/// Decode a Minimum Code Unit(MCU) as quickly as possible
///
/// # Arguments
/// - reader: The bitstream from where we read more bits.
/// - dc_table: The Huffman table used to decode the DC coefficient
/// - ac_table: The Huffman table used to decode AC values
/// - block: A memory region where we will write out the decoded values
/// - DC prediction: Last DC value for this component
///
#[allow(
clippy::many_single_char_names,
clippy::cast_possible_truncation,
clippy::cast_sign_loss
)]
#[inline(never)]
pub fn decode_mcu_block<T>(
&mut self, reader: &mut ZReader<T>, dc_table: &HuffmanTable, ac_table: &HuffmanTable,
qt_table: &[i32; DCT_BLOCK], block: &mut [i32; 64], dc_prediction: &mut i32
) -> Result<u16, DecodeErrors>
where
T: ZByteReaderTrait
{
// Get fast AC table as a reference before we enter the hot path
let ac_lookup = ac_table.ac_lookup.as_ref().unwrap();
let (mut symbol, mut r, mut fast_ac);
// Decode AC coefficients
let mut pos: usize = 1;
if self.bits_left < 1 && self.marker.is_some() {
return Err(DecodeErrors::Format(
"No more bytes left in stream before marker".to_string()
));
}
// decode DC, dc prediction will contain the value
self.decode_dc(reader, dc_table, dc_prediction)?;
// set dc to be the dc prediction.
block[0] = *dc_prediction * qt_table[0];
while pos < 64 {
self.refill(reader)?;
symbol = self.peek_bits::<HUFF_LOOKAHEAD>();
fast_ac = ac_lookup[symbol as usize];
symbol = ac_table.lookup[symbol as usize];
if fast_ac != 0 {
// FAST AC path
pos += ((fast_ac >> 4) & 15) as usize; // run
let t_pos = UN_ZIGZAG[min(pos, 63)] & 63;
block[t_pos] = i32::from(fast_ac >> 8) * (qt_table[t_pos]); // Value
self.drop_bits((fast_ac & 15) as u8);
pos += 1;
} else {
decode_huff!(self, symbol, ac_table);
r = symbol >> 4;
symbol &= 15;
if symbol != 0 {
pos += r as usize;
r = self.get_bits(symbol as u8);
symbol = huff_extend(r, symbol);
let t_pos = UN_ZIGZAG[pos & 63] & 63;
block[t_pos] = symbol * qt_table[t_pos];
pos += 1;
} else if r != 15 {
return Ok(pos as u16);
} else {
pos += 16;
}
}
}
return Ok(64);
}
/// Advance the bitstream over a block but ignore the data contained.
///
/// This updates DC prediction but we never dequantize and we never do any Zig-Zag translation
/// either. Still returns the index of the last component read.
pub fn discard_mcu_block<T>(
&mut self, reader: &mut ZReader<T>, dc_table: &HuffmanTable, ac_table: &HuffmanTable
) -> Result<u16, DecodeErrors>
where
T: ZByteReaderTrait
{
// Get fast AC table as a reference before we enter the hot path
let ac_lookup = ac_table.ac_lookup.as_ref().unwrap();
let (mut symbol, mut r, mut fast_ac);
// Decode AC coefficients
let mut pos: usize = 1;
// decode DC, dc prediction will contain the value
self.discard_dc(reader, dc_table)?;
while pos < 64 {
self.refill(reader)?;
symbol = self.peek_bits::<HUFF_LOOKAHEAD>();
fast_ac = ac_lookup[symbol as usize];
symbol = ac_table.lookup[symbol as usize];
if fast_ac != 0 {
// FAST AC path
pos += ((fast_ac >> 4) & 15) as usize; // run
self.drop_bits((fast_ac & 15) as u8);
pos += 1;
} else {
decode_huff!(self, symbol, ac_table);
r = symbol >> 4;
symbol &= 15;
if symbol != 0 {
pos += r as usize;
// Advance over bits but ignore.
let _ = self.get_bits(symbol as u8);
pos += 1;
} else if r != 15 {
return Ok(pos as u16);
} else {
pos += 16;
}
}
}
return Ok(64);
}
/// Peek `look_ahead` bits ahead without discarding them from the buffer
#[inline(always)]
#[allow(clippy::cast_possible_truncation)]
const fn peek_bits<const LOOKAHEAD: u8>(&self) -> i32 {
(self.aligned_buffer >> (64 - LOOKAHEAD)) as i32
}
/// Discard the next `N` bits without checking
#[inline]
fn drop_bits(&mut self, n: u8) {
// PS: Its a good check, but triggers fuzzer and a lot of false positives
//debug_assert!(self.bits_left >= n);
//self.bits_left -= n;
self.bits_left = self.bits_left.saturating_sub(n);
self.aligned_buffer <<= n;
}
/// Read `n_bits` from the buffer and discard them
#[inline(always)]
#[allow(clippy::cast_possible_truncation)]
fn get_bits(&mut self, n_bits: u8) -> i32 {
let mask = (1_u64 << n_bits) - 1;
self.aligned_buffer = self.aligned_buffer.rotate_left(u32::from(n_bits));
let bits = (self.aligned_buffer & mask) as i32;
self.bits_left = self.bits_left.wrapping_sub(n_bits);
bits
}
/// Decode a DC block
#[allow(clippy::cast_possible_truncation)]
#[inline]
pub(crate) fn decode_prog_dc_first<T>(
&mut self, reader: &mut ZReader<T>, dc_table: &HuffmanTable, block: &mut i16,
dc_prediction: &mut i32
) -> Result<(), DecodeErrors>
where
T: ZByteReaderTrait
{
self.decode_dc(reader, dc_table, dc_prediction)?;
*block = (*dc_prediction as i16).wrapping_mul(self.successive_low_mask);
return Ok(());
}
#[inline]
pub(crate) fn decode_prog_dc_refine<T>(
&mut self, reader: &mut ZReader<T>, block: &mut i16
) -> Result<(), DecodeErrors>
where
T: ZByteReaderTrait
{
// refinement scan
if self.bits_left < 1 {
self.refill(reader)?;
// if we find a marker, it may happens we don't refill.
// So let's confirm again that refill worked
if self.bits_left < 1 {
return Err(DecodeErrors::Format(
"Marker found where not expected in refine bit".to_string()
));
}
}
if self.get_bit() == 1 {
*block = block.wrapping_add(self.successive_low_mask);
}
Ok(())
}
/// Get a single bit from the bitstream
fn get_bit(&mut self) -> u8 {
let k = (self.aligned_buffer >> 63) as u8;
// discard a bit
self.drop_bits(1);
return k;
}
pub(crate) fn decode_mcu_ac_first<T>(
&mut self, reader: &mut ZReader<T>, ac_table: &HuffmanTable, block: &mut [i16; 64]
) -> Result<bool, DecodeErrors>
where
T: ZByteReaderTrait
{
let fast_ac = ac_table.ac_lookup.as_ref().unwrap();
let bit = self.successive_low_mask;
let mut k = self.spec_start as usize;
let (mut symbol, mut r, mut fac);
// EOB runs are handled in mcu_prog.rs
'block: loop {
self.refill(reader)?;
// Check for marker in the stream
symbol = self.peek_bits::<HUFF_LOOKAHEAD>();
fac = fast_ac[symbol as usize];
symbol = ac_table.lookup[symbol as usize];
if fac != 0 {
// fast ac path
k += ((fac >> 4) & 15) as usize; // run
block[UN_ZIGZAG[min(k, 63)] & 63] = (fac >> 8).wrapping_mul(bit); // value
self.drop_bits((fac & 15) as u8);
k += 1;
} else {
decode_huff!(self, symbol, ac_table);
r = symbol >> 4;
symbol &= 15;
if symbol != 0 {
k += r as usize;
r = self.get_bits(symbol as u8);
symbol = huff_extend(r, symbol);
block[UN_ZIGZAG[k & 63] & 63] = (symbol as i16).wrapping_mul(bit);
k += 1;
} else {
if r != 15 {
self.eob_run = 1 << r;
self.eob_run += self.get_bits(r as u8);
self.eob_run -= 1;
break;
}
k += 16;
}
}
if k > self.spec_end as usize {
break 'block;
}
}
return Ok(true);
}
#[allow(clippy::too_many_lines, clippy::op_ref)]
pub(crate) fn decode_mcu_ac_refine<T>(
&mut self, reader: &mut ZReader<T>, table: &HuffmanTable, block: &mut [i16; 64]
) -> Result<bool, DecodeErrors>
where
T: ZByteReaderTrait
{
let bit = self.successive_low_mask;
let mut k = self.spec_start;
let (mut symbol, mut r);
if self.eob_run == 0 {
'no_eob: loop {
// Decode a coefficient from the bit stream
self.refill(reader)?;
symbol = self.peek_bits::<HUFF_LOOKAHEAD>();
symbol = table.lookup[symbol as usize];
decode_huff!(self, symbol, table);
r = symbol >> 4;
symbol &= 15;
if symbol == 0 {
if r != 15 {
// EOB run is 2^r + bits
self.eob_run = 1 << r;
self.eob_run += self.get_bits(r as u8);
// EOB runs are handled by the eob logic
break 'no_eob;
}
} else {
if symbol != 1 {
return Err(DecodeErrors::HuffmanDecode(
"Bad Huffman code, corrupt JPEG?".to_string()
));
}
// get sign bit
// We assume we have enough bits, which should be correct for sane images
// since we refill by 32 above
if self.get_bit() == 1 {
symbol = i32::from(bit);
} else {
symbol = i32::from(-bit);
}
}
// Advance over already nonzero coefficients appending
// correction bits to the non-zeroes.
// A correction bit is 1 if the absolute value of the coefficient must be increased
if k <= self.spec_end {
'advance_nonzero: loop {
let coefficient = &mut block[UN_ZIGZAG[k as usize & 63] & 63];
if *coefficient != 0 {
if self.bits_left < 1 {
self.refill(reader)?;
if self.bits_left < 1 && self.marker.is_some() {
return Err(DecodeErrors::Format(
"Marker found where not expected in refine bit".to_string()
));
}
}
if self.get_bit() == 1 && (*coefficient & bit) == 0 {
if *coefficient > 0 {
*coefficient += bit;
} else {
*coefficient -= bit;
}
}
} else {
r -= 1;
if r < 0 {
// reached target zero coefficient.
break 'advance_nonzero;
}
};
if k == self.spec_end {
break 'advance_nonzero;
}
k += 1;
}
}
if symbol != 0 {
let pos = UN_ZIGZAG[k as usize & 63];
// output new non-zero coefficient.
block[pos & 63] = symbol as i16;
}
k += 1;
if k > self.spec_end {
break 'no_eob;
}
}
}
if self.eob_run > 0 {
// only run if block does not consists of purely zeroes
if &block[1..] != &[0; 63] {
self.refill(reader)?;
while k <= self.spec_end {
let coefficient = &mut block[UN_ZIGZAG[k as usize & 63] & 63];
if *coefficient != 0 && self.get_bit() == 1 {
// check if we already modified it, if so do nothing, otherwise
// append the correction bit.
if (*coefficient & bit) == 0 {
if *coefficient >= 0 {
*coefficient = coefficient.wrapping_add(bit);
} else {
*coefficient = coefficient.wrapping_sub(bit);
}
}
}
if self.bits_left < 1 {
// refill at the last possible moment
self.refill(reader)?;
}
k += 1;
}
}
// count a block completed in EOB run
self.eob_run -= 1;
}
return Ok(true);
}
pub fn update_progressive_params(&mut self, _ah: u8, al: u8, spec_start: u8, spec_end: u8) {
self.successive_low_mask = 1i16 << al;
self.spec_start = spec_start;
self.spec_end = spec_end;
}
/// Reset the stream if we have a restart marker
///
/// Restart markers indicate drop those bits in the stream and zero out
/// everything
#[cold]
pub fn reset(&mut self) {
self.bits_left = 0;
self.marker = None;
self.buffer = 0;
self.aligned_buffer = 0;
self.eob_run = 0;
}
}
/// Do the equivalent of JPEG HUFF_EXTEND
#[inline(always)]
fn huff_extend(x: i32, s: i32) -> i32 {
// if x<s return x else return x+offset[s] where offset[s] = ( (-1<<s)+1)
(x) + ((((x) - (1 << ((s) - 1))) >> 31) & (((-1) << (s)) + 1))
}
const fn has_zero(v: u32) -> bool {
// Retrieved from Stanford bithacks
// @ https://graphics.stanford.edu/~seander/bithacks.html#ZeroInWord
return !((((v & 0x7F7F_7F7F) + 0x7F7F_7F7F) | v) | 0x7F7F_7F7F) != 0;
}
const fn has_byte(b: u32, val: u8) -> bool {
// Retrieved from Stanford bithacks
// @ https://graphics.stanford.edu/~seander/bithacks.html#ZeroInWord
has_zero(b ^ ((!0_u32 / 255) * (val as u32)))
}
// mod tests {
// use zune_core::bytestream::ZCursor;
// use zune_core::colorspace::ColorSpace;
// use zune_core::options::DecoderOptions;
//
// use crate::JpegDecoder;
//
// #[test]
// fn test_image() {
// let img = "/Users/etemesi/Downloads/test_IDX_45_RAND_168601280367171438891916_minimized_837.jpg";
// let data = std::fs::read(img).unwrap();
// let options = DecoderOptions::new_cmd().jpeg_set_out_colorspace(ColorSpace::RGB);
// let mut decoder = JpegDecoder::new_with_options(ZCursor::new(&data[..]), options);
//
// decoder.decode().unwrap();
// println!("{:?}", decoder.options.jpeg_get_out_colorspace())
// }
// }
+102
View File
@@ -0,0 +1,102 @@
/*
* Copyright (c) 2023.
*
* This software is free software;
*
* You can redistribute it or modify it under terms of the MIT, Apache License or Zlib license
*/
#![allow(
clippy::many_single_char_names,
clippy::similar_names,
clippy::cast_possible_truncation,
clippy::cast_sign_loss,
clippy::cast_possible_wrap,
clippy::too_many_arguments,
clippy::doc_markdown
)]
//! Color space conversion routines
//!
//! This files exposes functions to convert one colorspace to another in a jpeg
//! image
//!
//! Currently supported conversions are
//!
//! - `YCbCr` to `RGB,RGBA,GRAYSCALE,RGBX`.
//!
//!
//! Hey there, if your reading this it means you probably need something, so let me help you.
//!
//! There are 3 supported cpu extensions here.
//! 1. Scalar
//! 2. SSE
//! 3. AVX
//!
//! There are two types of the color convert functions
//!
//! 1. Acts on 16 pixels.
//! 2. Acts on 8 pixels.
//!
//! The reason for this is because when implementing the AVX part it occurred to me that we can actually
//! do better and process 2 MCU's if we change IDCT return type to be `i16's`, since a lot of
//! CPU's these days support AVX extensions, it becomes nice if we optimize for that path ,
//! therefore AVX routines can process 16 pixels directly and SSE and Scalar just compensate.
//!
//! By compensating, I mean I wrote the 16 pixels version operating on the 8 pixel version twice.
//!
//! Therefore if your looking to optimize some routines, probably start there.
pub use scalar::ycbcr_to_grayscale;
use zune_core::colorspace::ColorSpace;
use zune_core::options::DecoderOptions;
#[cfg(any(target_arch = "x86", target_arch = "x86_64"))]
#[cfg(feature = "x86")]
pub use crate::color_convert::avx::{ycbcr_to_rgb_avx2, ycbcr_to_rgba_avx2};
use crate::decoder::ColorConvert16Ptr;
mod avx;
mod neon64;
mod scalar;
#[allow(unused_variables)]
pub fn choose_ycbcr_to_rgb_convert_func(
type_need: ColorSpace, options: &DecoderOptions
) -> Option<ColorConvert16Ptr> {
#[cfg(any(target_arch = "x86", target_arch = "x86_64"))]
#[cfg(feature = "x86")]
{
use zune_core::log::debug;
if options.use_avx2() {
debug!("Using AVX optimised color conversion functions");
// I believe avx2 means sse4 is also available
// match colorspace
match type_need {
ColorSpace::RGB => return Some(ycbcr_to_rgb_avx2),
ColorSpace::RGBA => return Some(ycbcr_to_rgba_avx2),
_ => () // fall through to scalar, which has more types
};
}
}
#[cfg(all(feature = "neon", target_arch = "aarch64"))]
{
if options.use_neon() {
use crate::color_convert::neon64::{ycbcr_to_rgb_neon, ycbcr_to_rgba_neon};
match type_need {
ColorSpace::RGB => return Some(ycbcr_to_rgb_neon),
ColorSpace::RGBA => return Some(ycbcr_to_rgba_neon),
_ => () // fall through to scalar, which has more types
};
}
}
// when there is no x86 or we haven't returned by here, resort to scalar
return match type_need {
ColorSpace::RGB => Some(scalar::ycbcr_to_rgb_inner_16_scalar::<false>),
ColorSpace::RGBA => Some(scalar::ycbcr_to_rgba_inner_16_scalar::<false>),
ColorSpace::BGRA => Some(scalar::ycbcr_to_rgba_inner_16_scalar::<true>),
ColorSpace::BGR => Some(scalar::ycbcr_to_rgb_inner_16_scalar::<true>),
_ => None
};
}
+297
View File
@@ -0,0 +1,297 @@
/*
* Copyright (c) 2023.
*
* This software is free software;
*
* You can redistribute it or modify it under terms of the MIT, Apache License or Zlib license
*/
//! AVX color conversion routines
//!
//! Okay these codes are cool
//!
//! Herein lies super optimized codes to do color conversions.
//!
//!
//! 1. The YCbCr to RGB use integer approximations and not the floating point equivalent.
//! That means we may be +- 2 of pixels generated by libjpeg-turbo jpeg decoding
//! (also libjpeg uses routines like `Y = 0.29900 * R + 0.33700 * G + 0.11400 * B + 0.25000 * G`)
//!
//! Firstly, we use integers (fun fact:there is no part of this code base where were dealing with
//! floating points.., fun fact: the first fun fact wasn't even fun.)
//!
//! Secondly ,we have cool clamping code, especially for rgba , where we don't need clamping and we
//! spend our time cursing that Intel decided permute instructions to work like 2 128 bit vectors(the compiler opitmizes
//! it out to something cool).
//!
//! There isn't a lot here (not as fun as bitstream ) but I hope you find what you're looking for.
//!
//! O and ~~subscribe to my youtube channel~~
#![cfg(any(target_arch = "x86", target_arch = "x86_64"))]
#![cfg(feature = "x86")]
#![allow(
clippy::wildcard_imports,
clippy::cast_possible_truncation,
clippy::too_many_arguments,
clippy::inline_always,
clippy::doc_markdown,
dead_code
)]
#[cfg(target_arch = "x86")]
use core::arch::x86::*;
#[cfg(target_arch = "x86_64")]
use core::arch::x86_64::*;
use crate::color_convert::scalar::{CB_CF, CR_CF, C_G_CB_COEF_2, C_G_CR_COEF_1, YUV_RND, Y_CF};
pub union YmmRegister {
// both are 32 when using std::mem::size_of
mm256: __m256i,
// for avx color conversion
array: [i16; 16]
}
const R_AVX_COEF: i32 = i32::from_ne_bytes([CR_CF.to_ne_bytes()[0], CR_CF.to_ne_bytes()[1], 0, 0]);
const B_AVX_COEF: i32 = i32::from_ne_bytes([0, 0, CB_CF.to_ne_bytes()[0], CB_CF.to_ne_bytes()[1]]);
const G_COEF_AVX_COEF: i32 = i32::from_ne_bytes([
C_G_CR_COEF_1.to_ne_bytes()[0],
C_G_CR_COEF_1.to_ne_bytes()[1],
C_G_CB_COEF_2.to_ne_bytes()[0],
C_G_CB_COEF_2.to_ne_bytes()[1]
]);
//--------------------------------------------------------------------------------------------------
// AVX conversion routines
//--------------------------------------------------------------------------------------------------
///
/// Convert YCBCR to RGB using AVX instructions
///
/// # Note
///**IT IS THE RESPONSIBILITY OF THE CALLER TO CALL THIS IN CPUS SUPPORTING
/// AVX2 OTHERWISE THIS IS UB**
///
/// *Peace*
///
/// This library itself will ensure that it's never called in CPU's not
/// supporting AVX2
///
/// # Arguments
/// - `y`,`cb`,`cr`: A reference of 8 i32's
/// - `out`: The output array where we store our converted items
/// - `offset`: The position from 0 where we write these RGB values
#[inline(always)]
pub fn ycbcr_to_rgb_avx2(
y: &[i16; 16], cb: &[i16; 16], cr: &[i16; 16], out: &mut [u8], offset: &mut usize
) {
// call this in another function to tell RUST to vectorize this
// storing
unsafe {
ycbcr_to_rgb_avx2_1(y, cb, cr, out, offset);
}
}
#[inline]
#[target_feature(enable = "avx2")]
unsafe fn ycbcr_to_rgb_avx2_1(
y: &[i16; 16], cb: &[i16; 16], cr: &[i16; 16], out: &mut [u8], offset: &mut usize
) {
let (mut r, mut g, mut b) = ycbcr_to_rgb_baseline_no_clamp(y, cb, cr);
r = _mm256_packus_epi16(r, _mm256_setzero_si256());
g = _mm256_packus_epi16(g, _mm256_setzero_si256());
b = _mm256_packus_epi16(b, _mm256_setzero_si256());
r = _mm256_permute4x64_epi64::<{ shuffle(3, 1, 2, 0) }>(r);
g = _mm256_permute4x64_epi64::<{ shuffle(3, 1, 2, 0) }>(g);
b = _mm256_permute4x64_epi64::<{ shuffle(3, 1, 2, 0) }>(b);
let sh_r = _mm256_setr_epi8(
0, 11, 6, 1, 12, 7, 2, 13, 8, 3, 14, 9, 4, 15, 10, 5, 0, 11, 6, 1, 12, 7, 2, 13, 8, 3, 14,
9, 4, 15, 10, 5
);
let sh_g = _mm256_setr_epi8(
5, 0, 11, 6, 1, 12, 7, 2, 13, 8, 3, 14, 9, 4, 15, 10, 5, 0, 11, 6, 1, 12, 7, 2, 13, 8, 3,
14, 9, 4, 15, 10
);
let sh_b = _mm256_setr_epi8(
10, 5, 0, 11, 6, 1, 12, 7, 2, 13, 8, 3, 14, 9, 4, 15, 10, 5, 0, 11, 6, 1, 12, 7, 2, 13, 8,
3, 14, 9, 4, 15
);
let r0 = _mm256_shuffle_epi8(r, sh_r);
let g0 = _mm256_shuffle_epi8(g, sh_g);
let b0 = _mm256_shuffle_epi8(b, sh_b);
let m0 = _mm256_setr_epi8(
0, -1, 0, 0, -1, 0, 0, -1, 0, 0, -1, 0, 0, -1, 0, 0, 0, -1, 0, 0, -1, 0, 0, -1, 0, 0, -1,
0, 0, -1, 0, 0
);
let m1 = _mm256_setr_epi8(
0, 0, -1, 0, 0, -1, 0, 0, -1, 0, 0, -1, 0, 0, -1, 0, 0, 0, -1, 0, 0, -1, 0, 0, -1, 0, 0,
-1, 0, 0, -1, 0
);
let p0 = _mm256_blendv_epi8(_mm256_blendv_epi8(r0, g0, m0), b0, m1);
let p1 = _mm256_blendv_epi8(_mm256_blendv_epi8(g0, b0, m0), r0, m1);
let p2 = _mm256_blendv_epi8(_mm256_blendv_epi8(b0, r0, m0), g0, m1);
let rgb0 = _mm256_permute2x128_si256::<32>(p0, p1);
let rgb1 = _mm256_permute2x128_si256::<48>(p2, p0);
_mm256_storeu_si256(out.as_mut_ptr().cast(), rgb0);
_mm_storeu_si128(out[32..].as_mut_ptr().cast(), _mm256_castsi256_si128(rgb1));
*offset += 48;
}
// Enabled avx2 automatically enables avx.
#[inline]
#[target_feature(enable = "avx2")]
/// A baseline implementation of YCbCr to RGB conversion which does not carry
/// out clamping
///
/// This is used by the `ycbcr_to_rgba_avx` and `ycbcr_to_rgbx` conversion
/// routines
unsafe fn ycbcr_to_rgb_baseline_no_clamp(
y: &[i16; 16], cb: &[i16; 16], cr: &[i16; 16]
) -> (__m256i, __m256i, __m256i) {
// Load values into a register
//
let y_c = _mm256_loadu_si256(y.as_ptr().cast());
let cb_c = _mm256_loadu_si256(cb.as_ptr().cast());
let cr_c = _mm256_loadu_si256(cr.as_ptr().cast());
// Here we want to use _mm256_madd_epi16 to perform 2 multiplications
// and one addition per instruction.
// At first, we have to pack i16 U and V that stores u8 into one u8 [U,V]
// then zero extend, and keep in mind that lanes is already been permuted.
let y_coeff = _mm256_set1_epi32(i32::from(Y_CF));
let cr_coeff = _mm256_set1_epi32(R_AVX_COEF);
let cb_coeff = _mm256_set1_epi32(B_AVX_COEF);
let cg_coeff = _mm256_set1_epi32(G_COEF_AVX_COEF);
let v_rnd = _mm256_set1_epi32(i32::from(YUV_RND));
let uv_bias = _mm256_set1_epi16(128);
// UV in memory because x86/x86_64 is always little endian
let v_0 = _mm256_slli_epi16::<8>(cb_c);
let u_v_8 = _mm256_or_si256(v_0, cr_c);
let mut u_v_lo = _mm256_unpacklo_epi8(u_v_8, _mm256_setzero_si256());
let mut u_v_hi = _mm256_unpackhi_epi8(u_v_8, _mm256_setzero_si256());
let mut y_lo = _mm256_unpacklo_epi16(y_c, _mm256_setzero_si256());
let mut y_hi = _mm256_unpackhi_epi16(y_c, _mm256_setzero_si256());
u_v_lo = _mm256_sub_epi16(u_v_lo, uv_bias);
u_v_hi = _mm256_sub_epi16(u_v_hi, uv_bias);
y_lo = _mm256_madd_epi16(y_lo, y_coeff);
y_hi = _mm256_madd_epi16(y_hi, y_coeff);
let mut r_lo = _mm256_madd_epi16(u_v_lo, cr_coeff);
let mut r_hi = _mm256_madd_epi16(u_v_hi, cr_coeff);
let mut g_lo = _mm256_madd_epi16(u_v_lo, cg_coeff);
let mut g_hi = _mm256_madd_epi16(u_v_hi, cg_coeff);
// This ordering is preferred to reduce register file pressure.
y_lo = _mm256_add_epi32(y_lo, v_rnd);
y_hi = _mm256_add_epi32(y_hi, v_rnd);
let mut b_lo = _mm256_madd_epi16(u_v_lo, cb_coeff);
let mut b_hi = _mm256_madd_epi16(u_v_hi, cb_coeff);
r_lo = _mm256_add_epi32(r_lo, y_lo);
r_hi = _mm256_add_epi32(r_hi, y_hi);
g_lo = _mm256_add_epi32(g_lo, y_lo);
g_hi = _mm256_add_epi32(g_hi, y_hi);
b_lo = _mm256_add_epi32(b_lo, y_lo);
b_hi = _mm256_add_epi32(b_hi, y_hi);
r_lo = _mm256_srai_epi32::<14>(r_lo);
r_hi = _mm256_srai_epi32::<14>(r_hi);
g_lo = _mm256_srai_epi32::<14>(g_lo);
g_hi = _mm256_srai_epi32::<14>(g_hi);
b_lo = _mm256_srai_epi32::<14>(b_lo);
b_hi = _mm256_srai_epi32::<14>(b_hi);
let r = _mm256_packus_epi32(r_lo, r_hi);
let g = _mm256_packus_epi32(g_lo, g_hi);
let b = _mm256_packus_epi32(b_lo, b_hi);
return (r, g, b);
}
#[inline(always)]
pub fn ycbcr_to_rgba_avx2(
y: &[i16; 16], cb: &[i16; 16], cr: &[i16; 16], out: &mut [u8], offset: &mut usize
) {
unsafe {
ycbcr_to_rgba_unsafe(y, cb, cr, out, offset);
}
}
#[inline]
#[target_feature(enable = "avx2")]
#[rustfmt::skip]
unsafe fn ycbcr_to_rgba_unsafe(
y: &[i16; 16], cb: &[i16; 16], cr: &[i16; 16],
out: &mut [u8],
offset: &mut usize,
)
{
// check if we have enough space to write.
let tmp:& mut [u8; 64] = out.get_mut(*offset..*offset + 64).expect("Slice to small cannot write").try_into().unwrap();
let (r, g, b) = ycbcr_to_rgb_baseline_no_clamp(y, cb, cr);
// set alpha channel to 255 for opaque
// And no these comments were not from me pressing the keyboard
// Pack the integers into u8's using unsigned saturation.
let c = _mm256_packus_epi16(r, g); //aaaaa_bbbbb_aaaaa_bbbbbb
let d = _mm256_packus_epi16(b, _mm256_set1_epi16(255)); // cccccc_dddddd_ccccccc_ddddd
// transpose_u16 and interleave channels
let e = _mm256_unpacklo_epi8(c, d); //ab_ab_ab_ab_ab_ab_ab_ab
let f = _mm256_unpackhi_epi8(c, d); //cd_cd_cd_cd_cd_cd_cd_cd
// final transpose_u16
let g = _mm256_unpacklo_epi8(e, f); //abcd_abcd_abcd_abcd_abcd
let h = _mm256_unpackhi_epi8(e, f);
// undo packus shuffling...
let i = _mm256_permute2x128_si256::<{ shuffle(3, 2, 1, 0) }>(g, h);
let j = _mm256_permute2x128_si256::<{ shuffle(1, 2, 3, 0) }>(g, h);
let k = _mm256_permute2x128_si256::<{ shuffle(3, 2, 0, 1) }>(g, h);
let l = _mm256_permute2x128_si256::<{ shuffle(0, 3, 2, 1) }>(g, h);
let m = _mm256_blend_epi32::<0b1111_0000>(i, j);
let n = _mm256_blend_epi32::<0b1111_0000>(k, l);
// Store
// Use streaming instructions to prevent polluting the cache?
_mm256_storeu_si256(tmp.as_mut_ptr().cast(), m);
_mm256_storeu_si256(tmp[32..].as_mut_ptr().cast(), n);
*offset += 64;
}
#[inline]
const fn shuffle(z: i32, y: i32, x: i32, w: i32) -> i32 {
(z << 6) | (y << 4) | (x << 2) | w
}
+144
View File
@@ -0,0 +1,144 @@
/*
* Copyright (c) 2025.
*
* This software is free software;
*
* You can redistribute it or modify it under terms of the MIT, Apache License or Zlib license
*/
//! Aarch64 color conversion routines
//! NEON is mandatory on aarch64.
#![cfg(all(feature = "neon", target_arch = "aarch64"))]
use core::arch::aarch64::*;
use crate::color_convert::scalar::{CB_CF, CR_CF, C_G_CB_COEF_2, C_G_CR_COEF_1, YUV_RND, Y_CF};
const C_1: u64 = u64::from_ne_bytes([
Y_CF.to_ne_bytes()[0],
Y_CF.to_ne_bytes()[1],
CR_CF.to_ne_bytes()[0],
CR_CF.to_ne_bytes()[1],
CB_CF.to_ne_bytes()[0],
CB_CF.to_ne_bytes()[1],
C_G_CR_COEF_1.to_ne_bytes()[0],
C_G_CR_COEF_1.to_ne_bytes()[1]
]);
const C_2: u64 = u64::from_ne_bytes([
C_G_CB_COEF_2.to_ne_bytes()[0],
C_G_CB_COEF_2.to_ne_bytes()[1],
0,
0,
0,
0,
0,
0
]);
#[inline(always)]
unsafe fn ycbcr_to_rgb_baseline_no_clamp(
y: &[i16; 16], cb: &[i16; 16], cr: &[i16; 16]
) -> (uint8x16_t, uint8x16_t, uint8x16_t) {
// NEON has 32 registers, so it is good idea to utilize a lot of variables at once
let cb_cr_bias = vdupq_n_s16(128);
// 0 - Y coeff, 1 - Cr, 2 - Cb, 3 - G1, 4 - G2
let coefficients = vcombine_s16(vcreate_s16(C_1), vcreate_s16(C_2));
let y0 = vld1q_s16(y.as_ptr().cast());
let y1 = vld1q_s16(y[8..].as_ptr().cast());
let mut cb0 = vld1q_s16(cb.as_ptr().cast());
let mut cb1 = vld1q_s16(cb[8..].as_ptr().cast());
let mut cr0 = vld1q_s16(cr.as_ptr().cast());
let mut cr1 = vld1q_s16(cr[8..].as_ptr().cast());
cb0 = vsubq_s16(cb0, cb_cr_bias);
cb1 = vsubq_s16(cb1, cb_cr_bias);
cr0 = vsubq_s16(cr0, cb_cr_bias);
cr1 = vsubq_s16(cr1, cb_cr_bias);
let bias = vdupq_n_s32(i32::from(YUV_RND));
let acc0 = vmlal_laneq_s16::<0>(bias, vget_low_s16(y0), coefficients);
let acc1 = vmlal_high_laneq_s16::<0>(bias, y0, coefficients);
let acc2 = vmlal_laneq_s16::<0>(bias, vget_low_s16(y1), coefficients);
let acc3 = vmlal_high_laneq_s16::<0>(bias, y1, coefficients);
let r0 = vmlal_laneq_s16::<1>(acc0, vget_low_s16(cr0), coefficients);
let r1 = vmlal_high_laneq_s16::<1>(acc1, cr0, coefficients);
let r2 = vmlal_laneq_s16::<1>(acc2, vget_low_s16(cr1), coefficients);
let r3 = vmlal_high_laneq_s16::<1>(acc3, cr1, coefficients);
let b0 = vmlal_laneq_s16::<2>(acc0, vget_low_s16(cb0), coefficients);
let b1 = vmlal_high_laneq_s16::<2>(acc1, cb0, coefficients);
let b2 = vmlal_laneq_s16::<2>(acc2, vget_low_s16(cb1), coefficients);
let b3 = vmlal_high_laneq_s16::<2>(acc3, cb1, coefficients);
// Saturating shift right with signed -> unsigned saturation
let qr0 = vqshrun_n_s32::<14>(r0);
let qr1 = vqshrun_n_s32::<14>(r1);
let qr2 = vqshrun_n_s32::<14>(r2);
let qr3 = vqshrun_n_s32::<14>(r3);
let mut g0 = vmlal_laneq_s16::<4>(acc0, vget_low_s16(cb0), coefficients);
let mut g1 = vmlal_high_laneq_s16::<4>(acc1, cb0, coefficients);
let mut g2 = vmlal_laneq_s16::<4>(acc2, vget_low_s16(cb1), coefficients);
let mut g3 = vmlal_high_laneq_s16::<4>(acc3, cb1, coefficients);
let qb0 = vqshrun_n_s32::<14>(b0);
let qb1 = vqshrun_n_s32::<14>(b1);
let qb2 = vqshrun_n_s32::<14>(b2);
let qb3 = vqshrun_n_s32::<14>(b3);
let r0 = vqmovn_u16(vcombine_u16(qr0, qr1));
let r1 = vqmovn_u16(vcombine_u16(qr2, qr3));
let b0 = vqmovn_u16(vcombine_u16(qb0, qb1));
let b1 = vqmovn_u16(vcombine_u16(qb2, qb3));
g0 = vmlal_laneq_s16::<3>(g0, vget_low_s16(cr0), coefficients);
g1 = vmlal_high_laneq_s16::<3>(g1, cr0, coefficients);
g2 = vmlal_laneq_s16::<3>(g2, vget_low_s16(cr1), coefficients);
g3 = vmlal_high_laneq_s16::<3>(g3, cr1, coefficients);
let qg0 = vqshrun_n_s32::<14>(g0);
let qg1 = vqshrun_n_s32::<14>(g1);
let qg2 = vqshrun_n_s32::<14>(g2);
let qg3 = vqshrun_n_s32::<14>(g3);
let g0 = vqmovn_u16(vcombine_u16(qg0, qg1));
let g1 = vqmovn_u16(vcombine_u16(qg2, qg3));
(
vcombine_u8(r0, r1),
vcombine_u8(g0, g1),
vcombine_u8(b0, b1)
)
}
#[inline(always)]
pub fn ycbcr_to_rgb_neon(
y: &[i16; 16], cb: &[i16; 16], cr: &[i16; 16], out: &mut [u8], offset: &mut usize
) {
// call this in another function to tell RUST to vectorize this
// storing
unsafe {
let (r, g, b) = ycbcr_to_rgb_baseline_no_clamp(y, cb, cr);
vst3q_u8(out.as_mut_ptr(), uint8x16x3_t(r, g, b));
*offset += 48;
}
}
#[inline(always)]
pub fn ycbcr_to_rgba_neon(
y: &[i16; 16], cb: &[i16; 16], cr: &[i16; 16], out: &mut [u8], offset: &mut usize
) {
unsafe {
let (r, g, b) = ycbcr_to_rgb_baseline_no_clamp(y, cb, cr);
vst4q_u8(out.as_mut_ptr(), uint8x16x4_t(r, g, b, vdupq_n_u8(255)));
*offset += 64;
}
}
+139
View File
@@ -0,0 +1,139 @@
/*
* Copyright (c) 2023.
*
* This software is free software;
*
* You can redistribute it or modify it under terms of the MIT, Apache License or Zlib license
*/
use core::convert::TryInto;
// Bt.601 Full Range inverse coefficients computed with 14 bits of precision with MPFR.
// This is important to keep them in i16.
// In most cases LLVM will detect what we're doing i16 widening to i32 math and will use
// appropriate optimizations.
pub(crate) const Y_CF: i16 = 16384;
pub(crate) const CR_CF: i16 = 22970;
pub(crate) const CB_CF: i16 = 29032;
pub(crate) const C_G_CR_COEF_1: i16 = -11700;
pub(crate) const C_G_CB_COEF_2: i16 = -5638;
pub(crate) const YUV_PREC: i16 = 14;
// Rounding const for YUV -> RGB conversion: floating equivalent 0.499(9).
pub(crate) const YUV_RND: i16 = (1 << (YUV_PREC - 1)) - 1;
/// Limit values to 0 and 255
#[inline]
#[allow(clippy::cast_possible_truncation, clippy::cast_sign_loss, dead_code)]
fn clamp(a: i32) -> u8 {
a.clamp(0, 255) as u8
}
/// YCbCr to RGBA color conversion
/// Convert YCbCr to RGB/BGR
///
/// Converts to RGB if const BGRA is false
///
/// Converts to BGR if const BGRA is true
pub fn ycbcr_to_rgba_inner_16_scalar<const BGRA: bool>(
y: &[i16; 16], cb: &[i16; 16], cr: &[i16; 16], output: &mut [u8], pos: &mut usize
) {
let (_, output_position) = output.split_at_mut(*pos);
// Convert into a slice with 64 elements for Rust to see we won't go out of bounds.
let opt: &mut [u8; 64] = output_position
.get_mut(0..64)
.expect("Slice to small cannot write")
.try_into()
.unwrap();
for ((&y, (cb, cr)), out) in y
.iter()
.zip(cb.iter().zip(cr.iter()))
.zip(opt.chunks_exact_mut(4))
{
let cr = cr - 128;
let cb = cb - 128;
let y0 = i32::from(y) * i32::from(Y_CF) + i32::from(YUV_RND);
let r = (y0 + i32::from(cr) * i32::from(CR_CF)) >> YUV_PREC;
let g = (y0
+ i32::from(cr) * i32::from(C_G_CR_COEF_1)
+ i32::from(cb) * i32::from(C_G_CB_COEF_2))
>> YUV_PREC;
let b = (y0 + i32::from(cb) * i32::from(CB_CF)) >> YUV_PREC;
if BGRA {
out[0] = clamp(b);
out[1] = clamp(g);
out[2] = clamp(r);
out[3] = 255;
} else {
out[0] = clamp(r);
out[1] = clamp(g);
out[2] = clamp(b);
out[3] = 255;
}
}
*pos += 64;
}
/// Convert YCbCr to RGB/BGR
///
/// Converts to RGB if const BGRA is false
///
/// Converts to BGR if const BGRA is true
pub fn ycbcr_to_rgb_inner_16_scalar<const BGRA: bool>(
y: &[i16; 16], cb: &[i16; 16], cr: &[i16; 16], output: &mut [u8], pos: &mut usize
) {
let (_, output_position) = output.split_at_mut(*pos);
// Convert into a slice with 48 elements
let opt: &mut [u8; 48] = output_position
.get_mut(0..48)
.expect("Slice to small cannot write")
.try_into()
.unwrap();
for ((&y, (cb, cr)), out) in y
.iter()
.zip(cb.iter().zip(cr.iter()))
.zip(opt.chunks_exact_mut(3))
{
let cr = cr - 128;
let cb = cb - 128;
let y0 = i32::from(y) * i32::from(Y_CF) + i32::from(YUV_RND);
let r = (y0 + i32::from(cr) * i32::from(CR_CF)) >> YUV_PREC;
let g = (y0
+ i32::from(cr) * i32::from(C_G_CR_COEF_1)
+ i32::from(cb) * i32::from(C_G_CB_COEF_2))
>> YUV_PREC;
let b = (y0 + i32::from(cb) * i32::from(CB_CF)) >> YUV_PREC;
if BGRA {
out[0] = clamp(b);
out[1] = clamp(g);
out[2] = clamp(r);
} else {
out[0] = clamp(r);
out[1] = clamp(g);
out[2] = clamp(b);
}
}
// Increment pos
*pos += 48;
}
pub fn ycbcr_to_grayscale(y: &[i16], width: usize, padded_width: usize, output: &mut [u8]) {
for (y_in, out) in y
.chunks_exact(padded_width)
.zip(output.chunks_exact_mut(width))
{
for (y, out) in y_in.iter().zip(out.iter_mut()) {
*out = *y as u8;
}
}
}
+232
View File
@@ -0,0 +1,232 @@
/*
* Copyright (c) 2023.
*
* This software is free software;
*
* You can redistribute it or modify it under terms of the MIT, Apache License or Zlib license
*/
//! This module exports a single struct to store information about
//! JPEG image components
//!
//! The data is extracted from a SOF header.
use alloc::vec::Vec;
use alloc::{format, vec};
use zune_core::log::trace;
use crate::alloc::string::ToString;
use crate::decoder::MAX_COMPONENTS;
use crate::errors::DecodeErrors;
use crate::upsampler::upsample_no_op;
const MAX_SAMP_FACTOR: usize = 4;
/// Represents an up-sampler function, this function will be called to upsample
/// a down-sampled image
pub type UpSampler = fn(
input: &[i16],
in_near: &[i16],
in_far: &[i16],
scratch_space: &mut [i16],
output: &mut [i16]
);
/// Component Data from start of frame
#[derive(Clone)]
pub(crate) struct Components {
/// The type of component that has the metadata below, can be Y,Cb or Cr
pub component_id: ComponentID,
/// Sub-sampling ratio of this component in the x-plane
pub vertical_sample: usize,
/// Sub-sampling ratio of this component in the y-plane
pub horizontal_sample: usize,
/// DC huffman table position
pub dc_huff_table: usize,
/// AC huffman table position for this element.
pub ac_huff_table: usize,
/// Quantization table number
pub quantization_table_number: u8,
/// Specifies quantization table to use with this component
pub quantization_table: [i32; 64],
/// dc prediction for the component
pub dc_pred: i32,
/// An up-sampling function, can be basic or SSE, depending
/// on the platform
pub up_sampler: UpSampler,
/// How pixels do we need to go to get to the next line?
pub width_stride: usize,
/// Component ID for progressive
pub id: u8,
/// Whether we need to decode this image component.
pub needed: bool,
/// Upsample scanline
pub raw_coeff: Vec<i16>,
/// Upsample destination, stores a scanline worth of sub sampled data
pub upsample_dest: Vec<i16>,
/// previous row, used to handle MCU boundaries
pub row_up: Vec<i16>,
/// current row, used to handle MCU boundaries again
pub row: Vec<i16>,
pub first_row_upsample_dest: Vec<i16>,
pub idct_pos: usize,
pub x: usize,
pub w2: usize,
pub y: usize,
pub sample_ratio: SampleRatios,
// a very annoying bug
pub fix_an_annoying_bug: usize
}
impl Components {
/// Create a new instance from three bytes from the start of frame
#[inline]
pub fn from(a: [u8; 3], pos: u8) -> Result<Components, DecodeErrors> {
// it's a unique identifier.
// doesn't have to be ascending
// see tests/inputs/huge_sof_number
//
// For such cases, use the position of the component
// to determine width
let id = match pos {
0 => ComponentID::Y,
1 => ComponentID::Cb,
2 => ComponentID::Cr,
3 => ComponentID::Q,
_ => {
return Err(DecodeErrors::Format(format!(
"Unknown component id found,{pos}, expected value between 1 and 4"
)))
}
};
let horizontal_sample = (a[1] >> 4) as usize;
let vertical_sample = (a[1] & 0x0f) as usize;
// Match libjpeg turbo on checking for sampling factors
// Reject anything above 4
if horizontal_sample > MAX_SAMP_FACTOR {
return Err(DecodeErrors::Format(format!(
"Bogus Horizontal Sampling Factor {horizontal_sample}"
)));
}
if vertical_sample > MAX_SAMP_FACTOR {
return Err(DecodeErrors::Format(format!(
"Bogus Vertical Sampling Factor {vertical_sample}"
)));
}
let quantization_table_number = a[2];
// confirm quantization number is between 0 and MAX_COMPONENTS
if usize::from(quantization_table_number) >= MAX_COMPONENTS {
return Err(DecodeErrors::Format(format!(
"Too large quantization number :{quantization_table_number}, expected value between 0 and {MAX_COMPONENTS}"
)));
}
// check that upsampling ratios are powers of two
// if these fail, it's probably a corrupt image.
if !horizontal_sample.is_power_of_two() {
return Err(DecodeErrors::Format(format!(
"Horizontal sample is not a power of two({horizontal_sample}) cannot decode"
)));
}
// if !vertical_sample.is_power_of_two() {
// return Err(DecodeErrors::Format(format!(
// "Vertical sub-sample is not power of two({vertical_sample}) cannot decode"
// )));
// }
if vertical_sample == 0 {
// Check for invalid vertical sample
return Err(DecodeErrors::Format("Vertical sample is zero".to_string()));
}
trace!(
"Component ID:{:?} \tHS:{} VS:{} QT:{}",
id,
horizontal_sample,
vertical_sample,
quantization_table_number
);
Ok(Components {
component_id: id,
vertical_sample,
horizontal_sample,
quantization_table_number,
first_row_upsample_dest: vec![],
// These two will be set with sof marker
dc_huff_table: 0,
ac_huff_table: 0,
quantization_table: [0; 64],
dc_pred: 0,
up_sampler: upsample_no_op,
// set later
width_stride: horizontal_sample,
id: a[0],
needed: true,
raw_coeff: vec![],
upsample_dest: vec![],
row_up: vec![],
row: vec![],
idct_pos: 0,
x: 0,
y: 0,
w2: 0,
sample_ratio: SampleRatios::None,
fix_an_annoying_bug: 1
})
}
/// Setup space for upsampling
///
/// During upsample, we need a reference of the last row so that upsampling can
/// proceed correctly,
/// so we store the last line of every scanline and use it for the next upsampling procedure
/// to store this, but since we don't need it for 1v1 upsampling,
/// we only call this for routines that need upsampling
///
/// # Requirements
/// - width stride of this element is set for the component.
pub fn setup_upsample_scanline(&mut self) {
self.row = vec![0; self.width_stride * self.vertical_sample];
self.row_up = vec![0; self.width_stride * self.vertical_sample];
self.first_row_upsample_dest =
vec![128; self.vertical_sample * self.width_stride * self.sample_ratio.sample()];
self.upsample_dest =
vec![0; self.width_stride * self.sample_ratio.sample() * self.fix_an_annoying_bug * 8];
}
}
/// Component ID's
#[derive(Copy, Debug, Clone, PartialEq, Eq)]
pub enum ComponentID {
/// Luminance channel
Y,
/// Blue chrominance
Cb,
/// Red chrominance
Cr,
/// Q or fourth component
Q
}
#[derive(Copy, Debug, Clone, PartialEq, Eq, Default)]
pub enum SampleRatios {
HV,
V,
H,
Generic(usize, usize),
#[default]
None
}
impl SampleRatios {
pub fn sample(self) -> usize {
match self {
SampleRatios::HV => 4,
SampleRatios::V | SampleRatios::H => 2,
SampleRatios::Generic(a, b) => a * b,
SampleRatios::None => 1
}
}
}
+987
View File
@@ -0,0 +1,987 @@
/*
* Copyright (c) 2023.
*
* This software is free software;
*
* You can redistribute it or modify it under terms of the MIT, Apache License or Zlib license
*/
//! Main image logic.
#![allow(clippy::doc_markdown)]
use alloc::string::ToString;
use alloc::vec::Vec;
use alloc::{format, vec};
use zune_core::bytestream::{ZByteReaderTrait, ZReader};
use zune_core::colorspace::ColorSpace;
use zune_core::log::{error, trace, warn};
use zune_core::options::DecoderOptions;
use crate::color_convert::choose_ycbcr_to_rgb_convert_func;
use crate::components::{Components, SampleRatios};
use crate::errors::{DecodeErrors, UnsupportedSchemes};
use crate::headers::{
parse_app1, parse_app13, parse_app14, parse_app2, parse_dqt, parse_huffman, parse_sos,
parse_start_of_frame
};
use crate::huffman::HuffmanTable;
use crate::idct::{choose_idct_func, choose_idct_1x1_func, choose_idct_4x4_func};
use crate::marker::Marker;
use crate::misc::SOFMarkers;
use crate::upsampler::{
choose_horizontal_samp_function, choose_hv_samp_function, choose_v_samp_function,
generic_sampler, upsample_no_op
};
/// Maximum components
pub(crate) const MAX_COMPONENTS: usize = 4;
/// Maximum image dimensions supported.
pub(crate) const MAX_DIMENSIONS: usize = 1 << 27;
/// Color conversion function that can convert YCbCr colorspace to RGB(A/X) for
/// 16 values
///
/// The following are guarantees to the following functions
///
/// 1. The `&[i16]` slices passed contain 16 items
///
/// 2. The slices passed are in the following order
/// `y,cb,cr`
///
/// 3. `&mut [u8]` is zero initialized
///
/// 4. `&mut usize` points to the position in the array where new values should
/// be used
///
/// The pointer should
/// 1. Carry out color conversion
/// 2. Update `&mut usize` with the new position
pub type ColorConvert16Ptr = fn(&[i16; 16], &[i16; 16], &[i16; 16], &mut [u8], &mut usize);
/// IDCT function prototype
///
/// This encapsulates a dequantize and IDCT function which will carry out the
/// following functions
///
/// Multiply each 64 element block of `&mut [i16]` with `&Aligned32<[i32;64]>`
/// Carry out IDCT (type 3 dct) on ach block of 64 i16's
pub type IDCTPtr = fn(&mut [i32; 64], &mut [i16], usize);
/// An encapsulation of an ICC chunk
pub(crate) struct ICCChunk {
pub(crate) seq_no: u8,
pub(crate) num_markers: u8,
pub(crate) data: Vec<u8>
}
/// A JPEG Decoder Instance.
#[allow(clippy::upper_case_acronyms, clippy::struct_excessive_bools)]
pub struct JpegDecoder<T> {
/// Struct to hold image information from SOI
pub(crate) info: ImageInfo,
/// Quantization tables, will be set to none and the tables will
/// be moved to `components` field
pub(crate) qt_tables: [Option<[i32; 64]>; MAX_COMPONENTS],
/// DC Huffman Tables with a maximum of 4 tables for each component
pub(crate) dc_huffman_tables: [Option<HuffmanTable>; MAX_COMPONENTS],
/// AC Huffman Tables with a maximum of 4 tables for each component
pub(crate) ac_huffman_tables: [Option<HuffmanTable>; MAX_COMPONENTS],
/// Image components, holds information like DC prediction and quantization
/// tables of a component
pub(crate) components: Vec<Components>,
/// maximum horizontal component of all channels in the image
pub(crate) h_max: usize,
// maximum vertical component of all channels in the image
pub(crate) v_max: usize,
/// mcu's width (interleaved scans)
pub(crate) mcu_width: usize,
/// MCU height(interleaved scans
pub(crate) mcu_height: usize,
/// Number of MCU's in the x plane
pub(crate) mcu_x: usize,
/// Number of MCU's in the y plane
pub(crate) mcu_y: usize,
/// Is the image interleaved?
pub(crate) is_interleaved: bool,
/// Image input colorspace, should be YCbCr for a sane image, might be
/// grayscale too
pub(crate) input_colorspace: ColorSpace,
// Progressive image details
/// Is the image progressive?
pub(crate) is_progressive: bool,
/// Start of spectral scan
pub(crate) spec_start: u8,
/// End of spectral scan
pub(crate) spec_end: u8,
/// Successive approximation bit position high
pub(crate) succ_high: u8,
/// Successive approximation bit position low
pub(crate) succ_low: u8,
/// Number of components.
pub(crate) num_scans: u8,
/// For a scan, check if any component has vertical/horizontal sampling.
pub(crate) scan_subsampled: bool,
// Function pointers, for pointy stuff.
/// Dequantize and idct function
// This is determined at runtime which function to run, statically it's
// initialized to a platform independent one and during initialization
// of this struct, we check if we can switch to a faster one which
// depend on certain CPU extensions.
pub(crate) idct_func: IDCTPtr,
/// Specialized IDCT when we can guarantee only few coefficients are non-zero.
///
/// **The callee must uphold a contract**. See [`choose_idct_4x4_func`].
pub(crate) idct_4x4_func: IDCTPtr,
pub(crate) idct_1x1_func: IDCTPtr,
// Color convert function which acts on 16 YCbCr values
pub(crate) color_convert_16: ColorConvert16Ptr,
pub(crate) z_order: [usize; MAX_COMPONENTS],
/// restart markers
pub(crate) restart_interval: usize,
pub(crate) todo: usize,
// decoder options
pub(crate) options: DecoderOptions,
// byte-stream
pub(crate) stream: ZReader<T>,
// Indicate whether headers have been decoded
pub(crate) headers_decoded: bool,
pub(crate) seen_sof: bool,
// exif data, lifted from app2
pub(crate) icc_data: Vec<ICCChunk>,
pub(crate) is_mjpeg: bool,
pub(crate) coeff: usize // Solves some weird bug :)
}
impl<T> JpegDecoder<T>
where
T: ZByteReaderTrait
{
#[allow(clippy::redundant_field_names)]
fn default(options: DecoderOptions, buffer: T) -> Self {
let color_convert = choose_ycbcr_to_rgb_convert_func(ColorSpace::RGB, &options).unwrap();
JpegDecoder {
info: ImageInfo::default(),
qt_tables: [None, None, None, None],
dc_huffman_tables: [None, None, None, None],
ac_huffman_tables: [None, None, None, None],
components: vec![],
// Interleaved information
h_max: 1,
v_max: 1,
mcu_height: 0,
mcu_width: 0,
mcu_x: 0,
mcu_y: 0,
is_interleaved: false,
is_progressive: false,
spec_start: 0,
spec_end: 0,
succ_high: 0,
succ_low: 0,
num_scans: 0,
scan_subsampled: false,
idct_func: choose_idct_func(&options),
idct_4x4_func: choose_idct_4x4_func(&options),
idct_1x1_func: choose_idct_1x1_func(&options),
color_convert_16: color_convert,
input_colorspace: ColorSpace::YCbCr,
z_order: [0; MAX_COMPONENTS],
restart_interval: 0,
todo: 0x7fff_ffff,
options: options,
stream: ZReader::new(buffer),
headers_decoded: false,
seen_sof: false,
icc_data: vec![],
is_mjpeg: false,
coeff: 1
}
}
/// Decode a buffer already in memory
///
/// The buffer should be a valid jpeg file, perhaps created by the command
/// `std:::fs::read()` or a JPEG file downloaded from the internet.
///
/// # Errors
/// See DecodeErrors for an explanation
pub fn decode(&mut self) -> Result<Vec<u8>, DecodeErrors> {
self.decode_headers()?;
let size = self.output_buffer_size().unwrap();
let mut out = vec![0; size];
self.decode_into(&mut out)?;
Ok(out)
}
/// Create a new Decoder instance
///
/// # Arguments
/// - `stream`: The raw bytes of a jpeg file.
#[must_use]
#[allow(clippy::new_without_default)]
pub fn new(stream: T) -> JpegDecoder<T> {
JpegDecoder::default(DecoderOptions::default(), stream)
}
/// Returns the image information
///
/// This **must** be called after a subsequent call to [`decode`] or [`decode_headers`]
/// it will return `None`
///
/// # Returns
/// - `Some(info)`: Image information,width, height, number of components
/// - None: Indicates image headers haven't been decoded
///
/// [`decode`]: JpegDecoder::decode
/// [`decode_headers`]: JpegDecoder::decode_headers
#[must_use]
pub fn info(&self) -> Option<ImageInfo> {
// we check for fails to that call by comparing what we have to the default, if
// it's default we assume that the caller failed to uphold the
// guarantees. We can be sure that an image cannot be the default since
// its a hard panic in-case width or height are set to zero.
if !self.headers_decoded {
return None;
}
return Some(self.info.clone());
}
/// Return the number of bytes required to hold a decoded image frame
/// decoded using the given input transformations
///
/// # Returns
/// - `Some(usize)`: Minimum size for a buffer needed to decode the image
/// - `None`: Indicates the image was not decoded, or image dimensions would overflow a usize
///
#[must_use]
pub fn output_buffer_size(&self) -> Option<usize> {
return if self.headers_decoded {
Some(
usize::from(self.width())
.checked_mul(usize::from(self.height()))?
.checked_mul(self.options.jpeg_get_out_colorspace().num_components())?
)
} else {
None
};
}
/// Get an immutable reference to the decoder options
/// for the decoder instance
///
/// This can be used to modify options before actual decoding
/// but after initial creation
///
/// # Example
/// ```no_run
/// use zune_core::bytestream::ZCursor;
/// use zune_jpeg::JpegDecoder;
///
/// let mut decoder = JpegDecoder::new(ZCursor::new(&[]));
/// // get current options
/// let mut options = decoder.options();
/// // modify it
/// let new_options = options.set_max_width(10);
/// // set it back
/// decoder.set_options(new_options);
///
/// ```
#[must_use]
pub const fn options(&self) -> &DecoderOptions {
&self.options
}
/// Return the input colorspace of the image
///
/// This indicates the colorspace that is present in
/// the image, but this may be different to the colorspace that
/// the output will be transformed to
///
/// # Returns
/// -`Some(Colorspace)`: Input colorspace
/// - None : Indicates the headers weren't decoded
#[must_use]
pub fn input_colorspace(&self) -> Option<ColorSpace> {
return if self.headers_decoded { Some(self.input_colorspace) } else { None };
}
/// Set decoder options
///
/// This can be used to set new options even after initialization
/// but before decoding.
///
/// This does not bear any significance after decoding an image
///
/// # Arguments
/// - `options`: New decoder options
///
/// # Example
/// Set maximum jpeg progressive passes to be 4
///
/// ```no_run
/// use zune_core::bytestream::ZCursor;
/// use zune_jpeg::JpegDecoder;
/// let mut decoder =JpegDecoder::new(ZCursor::new(&[]));
/// // this works also because DecoderOptions implements `Copy`
/// let options = decoder.options().jpeg_set_max_scans(4);
/// // set the new options
/// decoder.set_options(options);
/// // now decode
/// decoder.decode().unwrap();
/// ```
pub fn set_options(&mut self, options: DecoderOptions) {
self.options = options;
}
/// Decode Decoder headers
///
/// This routine takes care of parsing supported headers from a Decoder
/// image
///
/// # Supported Headers
/// - APP(0)
/// - SOF(O)
/// - DQT -> Quantization tables
/// - DHT -> Huffman tables
/// - SOS -> Start of Scan
/// # Unsupported Headers
/// - SOF(n) -> Decoder images which are not baseline/progressive
/// - DAC -> Images using Arithmetic tables
/// - JPG(n)
fn decode_headers_internal(&mut self) -> Result<(), DecodeErrors> {
if self.headers_decoded {
trace!("Headers decoded!");
return Ok(());
}
// match output colorspace here
// we know this will only be called once per image
// so makes sense
// We only care for ycbcr to rgb/rgba here
// in case one is using another colorspace.
// May god help you
let out_colorspace = self.options.jpeg_get_out_colorspace();
if matches!(
out_colorspace,
ColorSpace::BGR | ColorSpace::BGRA | ColorSpace::RGB | ColorSpace::RGBA
) {
self.color_convert_16 = choose_ycbcr_to_rgb_convert_func(
self.options.jpeg_get_out_colorspace(),
&self.options
)
.unwrap();
}
// First two bytes should be jpeg soi marker
let magic_bytes = self.stream.get_u16_be_err()?;
let mut last_byte = 0;
let mut bytes_before_marker = 0;
if magic_bytes != 0xffd8 {
return Err(DecodeErrors::IllegalMagicBytes(magic_bytes));
}
loop {
// read a byte
let mut m = self.stream.read_u8_err()?;
// AND OF COURSE some images will have fill bytes in their marker
// bitstreams because why not.
//
// I am disappointed as a man.
if (m == 0xFF || m == 0) && last_byte == 0xFF {
// This handles the edge case where
// images have markers with fill bytes(0xFF)
// or byte stuffing (0)
// I.e 0xFF 0xFF 0xDA
// and
// 0xFF 0 0xDA
// It should ignore those fill bytes and take 0xDA
// I don't know why such images exist
// but they do.
// so this is for you (with love)
while m == 0xFF || m == 0x0 {
last_byte = m;
m = self.stream.read_u8_err()?;
}
}
// Last byte should be 0xFF to confirm existence of a marker since markers look
// like OxFF(some marker data)
if last_byte == 0xFF {
let marker = Marker::from_u8(m);
if let Some(n) = marker {
if bytes_before_marker > 3 {
if self.options.strict_mode()
/*No reason to use this*/
{
return Err(DecodeErrors::FormatStatic(
"[strict-mode]: Extra bytes between headers"
));
}
error!(
"Extra bytes {} before marker 0xFF{:X}",
bytes_before_marker - 3,
m
);
}
bytes_before_marker = 0;
self.parse_marker_inner(n)?;
// break after reading the start of scan.
// what follows is the image data
if n == Marker::SOS {
self.headers_decoded = true;
trace!("Input colorspace {:?}", self.input_colorspace);
// Check if image is RGB
// The check is weird, we need to check if ID
// represents R, G and B in ascii,
//
// I am not sure if this is even specified in any standard,
// but jpegli https://github.com/google/jpegli does encode
// its images that way, so this will check for that. and handle it appropriately
// It is spefified here so that on a successful header decode,we can at least
// try to attribute image colorspace correctly.
//
// It was first the issue in https://github.com/etemesi254/zune-image/issues/291
// that brought it to light
//
let mut is_rgb = self.components.len() == 3;
let chars = ['R', 'G', 'B'];
for (comp, single_char) in self.components.iter().zip(chars.iter()) {
is_rgb &= comp.id == (*single_char) as u8
}
// Image is RGB, change colorspace
if is_rgb {
self.input_colorspace = ColorSpace::RGB;
}
return Ok(());
}
} else {
bytes_before_marker = 0;
warn!("Marker 0xFF{:X} not known", m);
let length = self.stream.get_u16_be_err()?;
if length < 2 {
return Err(DecodeErrors::Format(format!(
"Found a marker with invalid length : {length}"
)));
}
warn!("Skipping {} bytes", length - 2);
self.stream.skip((length - 2) as usize)?;
}
}
last_byte = m;
bytes_before_marker += 1;
}
// Check if image is RGB
}
#[allow(clippy::too_many_lines)]
pub(crate) fn parse_marker_inner(&mut self, m: Marker) -> Result<(), DecodeErrors> {
match m {
Marker::SOF(0..=2) => {
let marker = {
// choose marker
if m == Marker::SOF(0) || m == Marker::SOF(1) {
SOFMarkers::BaselineDct
} else {
self.is_progressive = true;
SOFMarkers::ProgressiveDctHuffman
}
};
trace!("Image encoding scheme =`{:?}`", marker);
// get components
parse_start_of_frame(marker, self)?;
}
// Start of Frame Segments not supported
Marker::SOF(v) => {
let feature = UnsupportedSchemes::from_int(v);
if let Some(feature) = feature {
return Err(DecodeErrors::Unsupported(feature));
}
return Err(DecodeErrors::Format("Unsupported image format".to_string()));
}
//APP(0) segment
Marker::APP(0) => {
let mut length = self.stream.get_u16_be_err()?;
if length < 2 {
return Err(DecodeErrors::Format(format!(
"Found a marker with invalid length:{length}\n"
)));
}
// skip for now
if length > 5 {
let mut buffer = [0u8; 5];
self.stream.read_exact_bytes(&mut buffer)?;
if &buffer == b"AVI1\0" {
self.is_mjpeg = true;
}
length -= 5;
}
self.stream.skip(length.saturating_sub(2) as usize)?;
//parse_app(buf, m, &mut self.info)?;
}
Marker::APP(1) => {
parse_app1(self)?;
}
Marker::APP(2) => {
parse_app2(self)?;
}
// Quantization tables
Marker::DQT => {
parse_dqt(self)?;
}
// Huffman tables
Marker::DHT => {
parse_huffman(self)?;
}
// Start of Scan Data
Marker::SOS => {
parse_sos(self)?;
}
Marker::EOI => return Err(DecodeErrors::FormatStatic("Premature End of image")),
Marker::DAC | Marker::DNL => {
return Err(DecodeErrors::Format(format!(
"Parsing of the following header `{m:?}` is not supported,\
cannot continue"
)));
}
Marker::DRI => {
if self.stream.get_u16_be_err()? != 4 {
return Err(DecodeErrors::Format(
"Bad DRI length, Corrupt JPEG".to_string()
));
}
self.restart_interval = usize::from(self.stream.get_u16_be_err()?);
trace!("DRI marker present ({})", self.restart_interval);
self.todo = self.restart_interval;
}
Marker::APP(14) => {
parse_app14(self)?;
}
Marker::APP(13) => {
parse_app13(self)?;
}
_ => {
warn!(
"Capabilities for processing marker \"{:?}\" not implemented",
m
);
let length = self.stream.get_u16_be_err()?;
if length < 2 {
return Err(DecodeErrors::Format(format!(
"Found a marker with invalid length:{length}\n"
)));
}
warn!("Skipping {} bytes", length - 2);
self.stream.skip((length - 2) as usize)?;
}
}
Ok(())
}
/// Get the embedded ICC profile if it exists
/// and is correct
///
/// One needs not to decode the whole image to extract this,
/// calling [`decode_headers`] for an image with an ICC profile
/// allows you to decode this
///
/// # Returns
/// - `Some(Vec<u8>)`: The raw ICC profile of the image
/// - `None`: May indicate an error in the ICC profile , non-existence of
/// an ICC profile, or that the headers weren't decoded.
///
/// [`decode_headers`]:Self::decode_headers
#[must_use]
pub fn icc_profile(&self) -> Option<Vec<u8>> {
let mut marker_present: [Option<&ICCChunk>; 256] = [None; 256];
if !self.headers_decoded {
return None;
}
let num_markers = self.icc_data.len();
if num_markers == 0 || num_markers >= 255 {
return None;
}
// check validity
for chunk in &self.icc_data {
if usize::from(chunk.num_markers) != num_markers {
// all the lengths must match
return None;
}
if chunk.seq_no == 0 {
warn!("Zero sequence number in ICC, corrupt ICC chunk");
return None;
}
if marker_present[usize::from(chunk.seq_no)].is_some() {
// duplicate seq_no
warn!("Duplicate sequence number in ICC, corrupt chunk");
return None;
}
marker_present[usize::from(chunk.seq_no)] = Some(chunk);
}
let mut data = Vec::with_capacity(1000);
// assemble the data now
for chunk in marker_present.get(1..=num_markers).unwrap() {
if let Some(ch) = chunk {
data.extend_from_slice(&ch.data);
} else {
warn!("Missing icc sequence number, corrupt ICC chunk ");
return None;
}
}
Some(data)
}
/// Return the exif data for the file
///
/// This returns the raw exif data starting at the
/// TIFF header
///
/// # Returns
/// -`Some(data)`: The raw exif data, if present in the image
/// - None: May indicate the following
///
/// 1. The image doesn't have exif data
/// 2. The image headers haven't been decoded
#[must_use]
pub fn exif(&self) -> Option<&Vec<u8>> {
return self.info.exif_data.as_ref();
}
/// Return the XMP data for the file
///
/// This returns raw XMP data starting at the XML header
/// One needs an XML/XMP decoder to extract valuable metadata
///
///
/// # Returns
/// - `Some(data)`: Raw xmp data
/// - `None`: May indicate the following
/// 1. The image does not have xmp data
/// 2. The image headers have not been decoded
///
/// # Example
///
/// ```no_run
/// use zune_core::bytestream::ZCursor;
/// use zune_jpeg::JpegDecoder;
/// let mut decoder = JpegDecoder::new(ZCursor::new(&[]));
/// // decode headers to extract xmp metadata if present
/// decoder.decode_headers().unwrap();
/// if let Some(data) = decoder.xmp(){
/// let stringified = String::from_utf8_lossy(data);
/// println!("XMP")
/// } else{
/// println!("No XMP Found")
/// }
///
/// ```
pub fn xmp(&self) -> Option<&Vec<u8>> {
return self.info.xmp_data.as_ref();
}
/// Return the IPTC data for the file
///
/// This returns the raw IPTC data.
///
/// # Returns
/// -`Some(data)`: The raw IPTC data, if present in the image
/// - None: May indicate the following
///
/// 1. The image doesn't have IPTC data
/// 2. The image headers haven't been decoded
#[must_use]
pub fn iptc(&self) -> Option<&Vec<u8>> {
return self.info.iptc_data.as_ref();
}
/// Get the output colorspace the image pixels will be decoded into
///
///
/// # Note.
/// This field can only be regarded after decoding headers,
/// as markers such as Adobe APP14 may dictate different colorspaces
/// than requested.
///
/// Calling `decode_headers` is sufficient to know what colorspace the
/// output is, if this is called after `decode` it indicates the colorspace
/// the output is currently in
///
/// Additionally not all input->output colorspace mappings are supported
/// but all input colorspaces can map to RGB colorspace, so that's a safe bet
/// if one is handling image formats
///
///# Returns
/// - `Some(Colorspace)`: If headers have been decoded, the colorspace the
///output array will be in
///- `None
#[must_use]
pub fn output_colorspace(&self) -> Option<ColorSpace> {
return if self.headers_decoded {
Some(self.options.jpeg_get_out_colorspace())
} else {
None
};
}
/// Decode into a pre-allocated buffer
///
/// It is an error if the buffer size is smaller than
/// [`output_buffer_size()`](Self::output_buffer_size)
///
/// If the buffer is bigger than expected, we ignore the end padding bytes
///
/// # Example
///
/// - Read headers and then alloc a buffer big enough to hold the image
///
/// ```no_run
/// use zune_core::bytestream::ZCursor;
/// use zune_jpeg::JpegDecoder;
/// let mut decoder = JpegDecoder::new(ZCursor::new(&[]));
/// // before we get output, we must decode the headers to get width
/// // height, and input colorspace
/// decoder.decode_headers().unwrap();
///
/// let mut out = vec![0;decoder.output_buffer_size().unwrap()];
/// // write into out
/// decoder.decode_into(&mut out).unwrap();
/// ```
///
///
pub fn decode_into(&mut self, out: &mut [u8]) -> Result<(), DecodeErrors> {
self.decode_headers_internal()?;
let expected_size = self.output_buffer_size().unwrap();
if out.len() < expected_size {
// too small of a size
return Err(DecodeErrors::TooSmallOutput(expected_size, out.len()));
}
// ensure we don't touch anyone else's scratch space
let out_len = core::cmp::min(out.len(), expected_size);
let out = &mut out[0..out_len];
if self.is_progressive {
self.decode_mcu_ycbcr_progressive(out)
} else {
self.decode_mcu_ycbcr_baseline(out)
}
}
/// Read only headers from a jpeg image buffer
///
/// This allows you to extract important information like
/// image width and height without decoding the full image
///
/// # Examples
/// ```no_run
/// use zune_core::bytestream::ZCursor;
/// use zune_jpeg::{JpegDecoder};
///
/// let img_data = std::fs::read("a_valid.jpeg").unwrap();
/// let mut decoder = JpegDecoder::new(ZCursor::new(&img_data));
/// decoder.decode_headers().unwrap();
///
/// println!("Total decoder dimensions are : {:?} pixels",decoder.dimensions());
/// println!("Number of components in the image are {}", decoder.info().unwrap().components);
/// ```
/// # Errors
/// See DecodeErrors enum for list of possible errors during decoding
pub fn decode_headers(&mut self) -> Result<(), DecodeErrors> {
self.decode_headers_internal()?;
Ok(())
}
/// Create a new decoder with the specified options to be used for decoding
/// an image
///
/// # Arguments
/// - `buf`: The input buffer from where we will pull in compressed jpeg bytes from
/// - `options`: Options specific to this decoder instance
#[must_use]
pub fn new_with_options(buf: T, options: DecoderOptions) -> JpegDecoder<T> {
JpegDecoder::default(options, buf)
}
/// Set up-sampling routines in case an image is down sampled
pub(crate) fn set_upsampling(&mut self) -> Result<(), DecodeErrors> {
// no sampling, return early
// check if horizontal max ==1
if self.h_max == self.v_max && self.h_max == 1 {
return Ok(());
}
for comp in &mut self.components {
let hs = self.h_max / comp.horizontal_sample;
let vs = self.v_max / comp.vertical_sample;
let samp_factor = match (hs, vs) {
(1, 1) => {
comp.sample_ratio = SampleRatios::None;
upsample_no_op
}
(2, 1) => {
comp.sample_ratio = SampleRatios::H;
choose_horizontal_samp_function(&self.options)
}
(1, 2) => {
comp.sample_ratio = SampleRatios::V;
choose_v_samp_function(&self.options)
}
(2, 2) => {
comp.sample_ratio = SampleRatios::HV;
choose_hv_samp_function(&self.options)
}
(hs, vs) => {
comp.sample_ratio = SampleRatios::Generic(hs, vs);
generic_sampler()
}
};
comp.setup_upsample_scanline();
comp.up_sampler = samp_factor;
}
return Ok(());
}
#[must_use]
/// Get the width of the image as a u16
///
/// The width lies between 1 and 65535
pub(crate) fn width(&self) -> u16 {
self.info.width
}
/// Get the height of the image as a u16
///
/// The height lies between 1 and 65535
#[must_use]
pub(crate) fn height(&self) -> u16 {
self.info.height
}
/// Get image dimensions as a tuple of width and height
/// or `None` if the image hasn't been decoded.
///
/// # Returns
/// - `Some(width,height)`: Image dimensions
/// - None : The image headers haven't been decoded
#[must_use]
pub const fn dimensions(&self) -> Option<(usize, usize)> {
return if self.headers_decoded {
Some((self.info.width as usize, self.info.height as usize))
} else {
None
};
}
}
#[derive(Default, Clone, Eq, PartialEq, Debug)]
pub struct GainMapInfo {
pub data: Vec<u8>
}
/// A struct representing Image Information
#[derive(Default, Clone, Eq, PartialEq)]
#[allow(clippy::module_name_repetitions)]
pub struct ImageInfo {
/// Width of the image
pub width: u16,
/// Height of image
pub height: u16,
/// PixelDensity
pub pixel_density: u8,
/// Start of frame markers
pub sof: SOFMarkers,
/// Horizontal sample
pub x_density: u16,
/// Vertical sample
pub y_density: u16,
/// Number of components
pub components: u8,
/// Gain Map information, useful for
/// UHDR images
pub gain_map_info: Vec<GainMapInfo>,
/// Multi picture information, useful for
/// UHDR images
pub multi_picture_information: Option<Vec<u8>>,
/// Exif Data
pub exif_data: Option<Vec<u8>>,
/// XMP Data
pub xmp_data: Option<Vec<u8>>,
/// IPTC Data
pub iptc_data: Option<Vec<u8>>,
/// Image sub-sampling ratio
pub sample_ratio: SampleRatios
}
impl ImageInfo {
/// Set width of the image
///
/// Found in the start of frame
pub(crate) fn set_width(&mut self, width: u16) {
self.width = width;
}
/// Set height of the image
///
/// Found in the start of frame
pub(crate) fn set_height(&mut self, height: u16) {
self.height = height;
}
/// Set the image density
///
/// Found in the start of frame
pub(crate) fn set_density(&mut self, density: u8) {
self.pixel_density = density;
}
/// Set image Start of frame marker
///
/// found in the Start of frame header
pub(crate) fn set_sof_marker(&mut self, marker: SOFMarkers) {
self.sof = marker;
}
/// Set image x-density(dots per pixel)
///
/// Found in the APP(0) marker
#[allow(dead_code)]
pub(crate) fn set_x(&mut self, sample: u16) {
self.x_density = sample;
}
/// Set image y-density
///
/// Found in the APP(0) marker
#[allow(dead_code)]
pub(crate) fn set_y(&mut self, sample: u16) {
self.y_density = sample;
}
}
+167
View File
@@ -0,0 +1,167 @@
/*
* Copyright (c) 2023.
*
* This software is free software;
*
* You can redistribute it or modify it under terms of the MIT, Apache License or Zlib license
*/
//! Contains most common errors that may be encountered in decoding a Decoder
//! image
use alloc::string::String;
use core::fmt::{Debug, Display, Formatter};
use zune_core::bytestream::ZByteIoError;
use crate::misc::{
START_OF_FRAME_EXT_AR, START_OF_FRAME_EXT_SEQ, START_OF_FRAME_LOS_SEQ,
START_OF_FRAME_LOS_SEQ_AR, START_OF_FRAME_PROG_DCT_AR
};
/// Common Decode errors
#[allow(clippy::module_name_repetitions)]
pub enum DecodeErrors {
/// Any other thing we do not know
Format(String),
/// Any other thing we do not know but we
/// don't need to allocate space on the heap
FormatStatic(&'static str),
/// Illegal Magic Bytes
IllegalMagicBytes(u16),
/// problems with the Huffman Tables in a Decoder file
HuffmanDecode(String),
/// Image has zero width
ZeroError,
/// Discrete Quantization Tables error
DqtError(String),
/// Start of scan errors
SosError(String),
/// Start of frame errors
SofError(String),
/// UnsupportedImages
Unsupported(UnsupportedSchemes),
/// MCU errors
MCUError(String),
/// Exhausted data
ExhaustedData,
/// Large image dimensions(Corrupted data)?
LargeDimensions(usize),
/// Too small output for size
TooSmallOutput(usize, usize),
IoErrors(ZByteIoError)
}
#[cfg(feature = "std")]
impl std::error::Error for DecodeErrors {}
impl From<&'static str> for DecodeErrors {
fn from(data: &'static str) -> Self {
return Self::FormatStatic(data);
}
}
impl From<ZByteIoError> for DecodeErrors {
fn from(data: ZByteIoError) -> Self {
return Self::IoErrors(data);
}
}
impl Debug for DecodeErrors {
fn fmt(&self, f: &mut Formatter<'_>) -> core::fmt::Result {
match &self
{
Self::Format(ref a) => write!(f, "{a:?}"),
Self::FormatStatic(a) => write!(f, "{:?}", &a),
Self::HuffmanDecode(ref reason) =>
{
write!(f, "Error decoding huffman values: {reason}")
}
Self::ZeroError => write!(f, "Image width or height is set to zero, cannot continue"),
Self::DqtError(ref reason) => write!(f, "Error parsing DQT segment. Reason:{reason}"),
Self::SosError(ref reason) => write!(f, "Error parsing SOS Segment. Reason:{reason}"),
Self::SofError(ref reason) => write!(f, "Error parsing SOF segment. Reason:{reason}"),
Self::IllegalMagicBytes(bytes) =>
{
write!(f, "Error parsing image. Illegal start bytes:{bytes:X}")
}
Self::MCUError(ref reason) => write!(f, "Error in decoding MCU. Reason {reason}"),
Self::Unsupported(ref image_type) =>
{
write!(f, "{image_type:?}")
}
Self::ExhaustedData => write!(f, "Exhausted data in the image"),
Self::LargeDimensions(ref dimensions) => write!(
f,
"Too large dimensions {dimensions},library supports up to {}", crate::decoder::MAX_DIMENSIONS
),
Self::TooSmallOutput(expected, found) => write!(f, "Too small output, expected buffer with at least {expected} bytes but got one with {found} bytes"),
Self::IoErrors(error)=>write!(f,"I/O errors {error:?}"),
}
}
}
impl Display for DecodeErrors {
fn fmt(&self, f: &mut Formatter<'_>) -> core::fmt::Result {
write!(f, "{self:?}")
}
}
/// Contains Unsupported/Yet-to-be supported Decoder image encoding types.
#[derive(Eq, PartialEq, Copy, Clone)]
pub enum UnsupportedSchemes {
/// SOF_1 Extended sequential DCT,Huffman coding
ExtendedSequentialHuffman,
/// Lossless (sequential), huffman coding,
LosslessHuffman,
/// Extended sequential DEC, arithmetic coding
ExtendedSequentialDctArithmetic,
/// Progressive DCT, arithmetic coding,
ProgressiveDctArithmetic,
/// Lossless ( sequential), arithmetic coding
LosslessArithmetic
}
impl Debug for UnsupportedSchemes {
fn fmt(&self, f: &mut Formatter<'_>) -> core::fmt::Result {
match &self {
Self::ExtendedSequentialHuffman => {
write!(f, "The library cannot yet decode images encoded using Extended Sequential Huffman encoding scheme yet.")
}
Self::LosslessHuffman => {
write!(f, "The library cannot yet decode images encoded with Lossless Huffman encoding scheme")
}
Self::ExtendedSequentialDctArithmetic => {
write!(f,"The library cannot yet decode Images Encoded with Extended Sequential DCT Arithmetic scheme")
}
Self::ProgressiveDctArithmetic => {
write!(f,"The library cannot yet decode images encoded with Progressive DCT Arithmetic scheme")
}
Self::LosslessArithmetic => {
write!(f,"The library cannot yet decode images encoded with Lossless Arithmetic encoding scheme")
}
}
}
}
impl UnsupportedSchemes {
#[must_use]
/// Create an unsupported scheme from an integer
///
/// # Returns
/// `Some(UnsupportedScheme)` if the int refers to a specific scheme,
/// otherwise returns `None`
pub fn from_int(int: u8) -> Option<UnsupportedSchemes> {
let int = u16::from_be_bytes([0xff, int]);
match int {
START_OF_FRAME_PROG_DCT_AR => Some(Self::ProgressiveDctArithmetic),
START_OF_FRAME_LOS_SEQ => Some(Self::LosslessHuffman),
START_OF_FRAME_LOS_SEQ_AR => Some(Self::LosslessArithmetic),
START_OF_FRAME_EXT_SEQ => Some(Self::ExtendedSequentialHuffman),
START_OF_FRAME_EXT_AR => Some(Self::ExtendedSequentialDctArithmetic),
_ => None
}
}
}
+662
View File
@@ -0,0 +1,662 @@
/*
* Copyright (c) 2023.
*
* This software is free software;
*
* You can redistribute it or modify it under terms of the MIT, Apache License or Zlib license
*/
//! Decode Decoder markers/segments
//!
//! This file deals with decoding header information in a jpeg file
//!
use alloc::format;
use alloc::string::ToString;
use alloc::vec::Vec;
use zune_core::bytestream::ZByteReaderTrait;
use zune_core::colorspace::ColorSpace;
use zune_core::log::{debug, trace, warn};
use core::cmp::max;
use crate::components::{Components, SampleRatios};
use crate::decoder::{GainMapInfo, ICCChunk, JpegDecoder, MAX_COMPONENTS};
use crate::errors::DecodeErrors;
use crate::huffman::HuffmanTable;
use crate::misc::{SOFMarkers, UN_ZIGZAG};
///**B.2.4.2 Huffman table-specification syntax**
#[allow(clippy::similar_names, clippy::cast_sign_loss)]
pub(crate) fn parse_huffman<T: ZByteReaderTrait>(
decoder: &mut JpegDecoder<T>
) -> Result<(), DecodeErrors>
where
{
// Read the length of the Huffman table
let mut dht_length = i32::from(decoder.stream.get_u16_be_err()?.checked_sub(2).ok_or(
DecodeErrors::FormatStatic("Invalid Huffman length in image")
)?);
while dht_length > 16 {
// HT information
let ht_info = decoder.stream.read_u8_err()?;
// third bit indicates whether the huffman encoding is DC or AC type
let dc_or_ac = (ht_info >> 4) & 0xF;
// Indicate the position of this table, should be less than 4;
let index = (ht_info & 0xF) as usize;
// read the number of symbols
let mut num_symbols: [u8; 17] = [0; 17];
if index >= MAX_COMPONENTS {
return Err(DecodeErrors::HuffmanDecode(format!(
"Invalid DHT index {index}, expected between 0 and 3"
)));
}
if dc_or_ac > 1 {
return Err(DecodeErrors::HuffmanDecode(format!(
"Invalid DHT position {dc_or_ac}, should be 0 or 1"
)));
}
decoder.stream.read_exact_bytes(&mut num_symbols[1..17])?;
dht_length -= 1 + 16;
let symbols_sum: i32 = num_symbols.iter().map(|f| i32::from(*f)).sum();
// The sum of the number of symbols cannot be greater than 256;
if symbols_sum > 256 {
return Err(DecodeErrors::FormatStatic(
"Encountered Huffman table with excessive length in DHT"
));
}
if symbols_sum > dht_length {
return Err(DecodeErrors::HuffmanDecode(format!(
"Excessive Huffman table of length {symbols_sum} found when header length is {dht_length}"
)));
}
dht_length -= symbols_sum;
// A table containing symbols in increasing code length
let mut symbols = [0; 256];
decoder
.stream
.read_exact_bytes(&mut symbols[0..(symbols_sum as usize)])?;
// store
match dc_or_ac {
0 => {
decoder.dc_huffman_tables[index] = Some(HuffmanTable::new(
&num_symbols,
symbols,
true,
decoder.is_progressive
)?);
}
_ => {
decoder.ac_huffman_tables[index] = Some(HuffmanTable::new(
&num_symbols,
symbols,
false,
decoder.is_progressive
)?);
}
}
}
if dht_length > 0 {
return Err(DecodeErrors::FormatStatic("Bogus Huffman table definition"));
}
Ok(())
}
///**B.2.4.1 Quantization table-specification syntax**
#[allow(clippy::cast_possible_truncation, clippy::needless_range_loop)]
pub(crate) fn parse_dqt<T: ZByteReaderTrait>(img: &mut JpegDecoder<T>) -> Result<(), DecodeErrors> {
// read length
let mut qt_length =
img.stream
.get_u16_be_err()?
.checked_sub(2)
.ok_or(DecodeErrors::FormatStatic(
"Invalid DQT length. Length should be greater than 2"
))?;
// A single DQT header may have multiple QT's
while qt_length > 0 {
let qt_info = img.stream.read_u8_err()?;
// 0 = 8 bit otherwise 16 bit dqt
let precision = (qt_info >> 4) as usize;
// last 4 bits give us position
let table_position = (qt_info & 0x0f) as usize;
let precision_value = 64 * (precision + 1);
if (precision_value + 1) as u16 > qt_length {
return Err(DecodeErrors::DqtError(format!("Invalid QT table bytes left :{}. Too small to construct a valid qt table which should be {} long", qt_length, precision_value + 1)));
}
let dct_table = match precision {
0 => {
let mut qt_values = [0; 64];
img.stream.read_exact_bytes(&mut qt_values)?;
qt_length -= (precision_value as u16) + 1 /*QT BIT*/;
// carry out un zig-zag here
un_zig_zag(&qt_values)
}
1 => {
// 16 bit quantization tables
let mut qt_values = [0_u16; 64];
for i in 0..64 {
qt_values[i] = img.stream.get_u16_be_err()?;
}
qt_length -= (precision_value as u16) + 1;
un_zig_zag(&qt_values)
}
_ => {
return Err(DecodeErrors::DqtError(format!(
"Expected QT precision value of either 0 or 1, found {precision:?}"
)));
}
};
if table_position >= MAX_COMPONENTS {
return Err(DecodeErrors::DqtError(format!(
"Too large table position for QT :{table_position}, expected between 0 and 3"
)));
}
trace!("Assigning qt table {table_position} with precision {precision}");
img.qt_tables[table_position] = Some(dct_table);
}
return Ok(());
}
/// Section:`B.2.2 Frame header syntax`
pub(crate) fn parse_start_of_frame<T: ZByteReaderTrait>(
sof: SOFMarkers, img: &mut JpegDecoder<T>
) -> Result<(), DecodeErrors> {
if img.seen_sof {
return Err(DecodeErrors::SofError(
"Two Start of Frame Markers".to_string()
));
}
// Get length of the frame header
let length = img.stream.get_u16_be_err()?;
// usually 8, but can be 12 and 16, we currently support only 8
// so sorry about that 12 bit images
let dt_precision = img.stream.read_u8_err()?;
if dt_precision != 8 {
return Err(DecodeErrors::SofError(format!(
"The library can only parse 8-bit images, the image has {dt_precision} bits of precision"
)));
}
img.info.set_density(dt_precision);
// read and set the image height.
let img_height = img.stream.get_u16_be_err()?;
img.info.set_height(img_height);
// read and set the image width
let img_width = img.stream.get_u16_be_err()?;
img.info.set_width(img_width);
trace!("Image width :{}", img_width);
trace!("Image height :{}", img_height);
if usize::from(img_width) > img.options.max_width() {
return Err(DecodeErrors::Format(format!("Image width {} greater than width limit {}. If use `set_limits` if you want to support huge images", img_width, img.options.max_width())));
}
if usize::from(img_height) > img.options.max_height() {
return Err(DecodeErrors::Format(format!("Image height {} greater than height limit {}. If use `set_limits` if you want to support huge images", img_height, img.options.max_height())));
}
// Check image width or height is zero
if img_width == 0 || img_height == 0 {
return Err(DecodeErrors::ZeroError);
}
// Number of components for the image.
let num_components = img.stream.read_u8_err()?;
if num_components == 0 {
return Err(DecodeErrors::SofError(
"Number of components cannot be zero.".to_string()
));
}
let expected = 8 + 3 * u16::from(num_components);
// length should be equal to num components
if length != expected {
return Err(DecodeErrors::SofError(format!(
"Length of start of frame differs from expected {expected},value is {length}"
)));
}
trace!("Image components : {}", num_components);
if num_components == 1 {
// SOF sets the number of image components
// and that to us translates to setting input and output
// colorspaces to zero
img.input_colorspace = ColorSpace::Luma;
//img.options = img.options.jpeg_set_out_colorspace(ColorSpace::Luma);
debug!("Overriding default colorspace set to Luma");
}
if num_components == 4 && img.input_colorspace == ColorSpace::YCbCr {
trace!("Input image has 4 components, defaulting to CMYK colorspace");
// https://entropymine.wordpress.com/2018/10/22/how-is-a-jpeg-images-color-type-determined/
img.input_colorspace = ColorSpace::CMYK;
}
// set number of components
img.info.components = num_components;
let mut components = Vec::with_capacity(num_components as usize);
let mut temp = [0; 3];
for pos in 0..num_components {
// read 3 bytes for each component
img.stream.read_exact_bytes(&mut temp)?;
// create a component.
let component = Components::from(temp, pos)?;
components.push(component);
}
img.seen_sof = true;
img.info.set_sof_marker(sof);
img.components = components;
let mut h_max = 1;
let mut v_max = 1;
for comp in &img.components {
h_max = max(h_max, comp.horizontal_sample);
v_max = max(v_max, comp.vertical_sample);
}
img.info.sample_ratio = match (h_max, v_max) {
(1, 1) => SampleRatios::None,
(1, 2) => SampleRatios::V,
(2, 1) => SampleRatios::H,
(2, 2) => SampleRatios::HV,
(hs, vs) => SampleRatios::Generic(hs, vs)
};
Ok(())
}
/// Parse a start of scan data
pub(crate) fn parse_sos<T: ZByteReaderTrait>(
image: &mut JpegDecoder<T>
) -> Result<(), DecodeErrors> {
// Scan header length
let ls = usize::from(image.stream.get_u16_be_err()?);
// Number of image components in scan
let ns = image.stream.read_u8_err()?;
let mut seen: [_; 5] = [-1; { MAX_COMPONENTS + 1 }];
image.num_scans = ns;
let smallest_size = 6 + 2 * usize::from(ns);
if ls != smallest_size {
return Err(DecodeErrors::SosError(format!(
"Bad SOS length {ls},corrupt jpeg"
)));
}
// Check number of components.
if !(1..5).contains(&ns) {
return Err(DecodeErrors::SosError(format!(
"Invalid number of components in start of scan {ns}, expected in range 1..5"
)));
}
if image.info.components == 0 {
return Err(DecodeErrors::FormatStatic(
"Error decoding SOF Marker, Number of components cannot be zero."
));
}
// consume spec parameters
image.scan_subsampled = false;
for i in 0..ns {
let id = image.stream.read_u8_err()?;
if seen.contains(&i32::from(id)) {
return Err(DecodeErrors::SofError(format!(
"Duplicate ID {id} seen twice in the same component"
)));
}
seen[usize::from(i)] = i32::from(id);
// DC and AC huffman table position
// top 4 bits contain dc huffman destination table
// lower four bits contain ac huffman destination table
let y = image.stream.read_u8_err()?;
let mut j = 0;
while j < image.info.components {
if image.components[j as usize].id == id {
break;
}
j += 1;
}
if j == image.info.components {
return Err(DecodeErrors::SofError(format!(
"Invalid component id {}, expected one one of {:?}",
id,
image.components.iter().map(|c| c.id).collect::<Vec<_>>()
)));
}
let component = &mut image.components[usize::from(j)];
component.dc_huff_table = usize::from((y >> 4) & 0xF);
component.ac_huff_table = usize::from(y & 0xF);
image.z_order[i as usize] = j as usize;
if component.vertical_sample != 1 || component.horizontal_sample != 1 {
image.scan_subsampled = true;
}
trace!(
"Assigned huffman tables {}/{} to component {j}, id={}",
image.components[usize::from(j)].dc_huff_table,
image.components[usize::from(j)].ac_huff_table,
image.components[usize::from(j)].id,
);
}
// Collect the component spec parameters
// This is only needed for progressive images but I'll read
// them in order to ensure they are correct according to the spec
// Extract progressive information
// https://www.w3.org/Graphics/JPEG/itu-t81.pdf
// Page 42
// Start of spectral / predictor selection. (between 0 and 63)
image.spec_start = image.stream.read_u8_err()?;
// End of spectral selection
image.spec_end = image.stream.read_u8_err()?;
let bit_approx = image.stream.read_u8_err()?;
// successive approximation bit position high
image.succ_high = bit_approx >> 4;
if image.spec_end > 63 {
return Err(DecodeErrors::SosError(format!(
"Invalid Se parameter {}, range should be 0-63",
image.spec_end
)));
}
if image.spec_start > 63 {
return Err(DecodeErrors::SosError(format!(
"Invalid Ss parameter {}, range should be 0-63",
image.spec_start
)));
}
if image.succ_high > 13 {
return Err(DecodeErrors::SosError(format!(
"Invalid Ah parameter {}, range should be 0-13",
image.succ_low
)));
}
// successive approximation bit position low
image.succ_low = bit_approx & 0xF;
if image.succ_low > 13 {
return Err(DecodeErrors::SosError(format!(
"Invalid Al parameter {}, range should be 0-13",
image.succ_low
)));
}
// skip any bytes not read
image.stream.skip(smallest_size.saturating_sub(ls))?;
trace!(
"Ss={}, Se={} Ah={} Al={}",
image.spec_start,
image.spec_end,
image.succ_high,
image.succ_low
);
Ok(())
}
/// Parse the APP13 (IPTC) segment.
pub(crate) fn parse_app13<T: ZByteReaderTrait>(
decoder: &mut JpegDecoder<T>
) -> Result<(), DecodeErrors> {
const IPTC_PREFIX: &[u8] = b"Photoshop 3.0\0";
// skip length.
let mut length = usize::from(decoder.stream.get_u16_be());
if length < 2 {
return Err(DecodeErrors::FormatStatic("Too small APP13 length"));
}
// length bytes.
length -= 2;
if length > IPTC_PREFIX.len() && decoder.stream.peek_at(0, IPTC_PREFIX.len())? == IPTC_PREFIX {
// skip bytes we read above.
decoder.stream.skip(IPTC_PREFIX.len())?;
length -= IPTC_PREFIX.len();
let iptc_bytes = decoder.stream.peek_at(0, length)?.to_vec();
decoder.info.iptc_data = Some(iptc_bytes);
}
decoder.stream.skip(length)?;
Ok(())
}
/// Parse Adobe App14 segment
pub(crate) fn parse_app14<T: ZByteReaderTrait>(
decoder: &mut JpegDecoder<T>
) -> Result<(), DecodeErrors> {
// skip length
let mut length = usize::from(decoder.stream.get_u16_be());
if length < 2 {
return Err(DecodeErrors::FormatStatic("Too small APP14 length"));
}
if decoder.stream.peek_at(0, 5)? == b"Adobe" {
if length < 14 {
return Err(DecodeErrors::FormatStatic(
"Too short of a length for App14 segment"
));
}
// move stream 6 bytes to remove adobe id
decoder.stream.skip(6)?;
// skip version, flags0 and flags1
decoder.stream.skip(5)?;
// get color transform
let transform = decoder.stream.read_u8();
// https://exiftool.org/TagNames/JPEG.html#Adobe
match transform {
0 => decoder.input_colorspace = ColorSpace::CMYK,
1 => decoder.input_colorspace = ColorSpace::YCbCr,
2 => decoder.input_colorspace = ColorSpace::YCCK,
_ => {
return Err(DecodeErrors::Format(format!(
"Unknown Adobe colorspace {transform}"
)))
}
}
// length = 2
// adobe id = 6
// version = 5
// transform = 1
length = length.saturating_sub(14);
} else {
warn!("Not a valid Adobe APP14 Segment, skipping {} bytes", length);
length = length.saturating_sub(2);
}
// skip any proceeding lengths.
// we do not need them
decoder.stream.skip(length)?;
Ok(())
}
/// Parse the APP1 segment
///
/// This contains the exif tag
pub(crate) fn parse_app1<T: ZByteReaderTrait>(
decoder: &mut JpegDecoder<T>
) -> Result<(), DecodeErrors> {
const XMP_NAMESPACE_PREFIX: &[u8] = b"http://ns.adobe.com/xap/1.0/\0";
// contains exif data
let mut length = usize::from(decoder.stream.get_u16_be());
if length < 2 {
return Err(DecodeErrors::FormatStatic("Too small app1 length"));
}
// length bytes
length -= 2;
if length > 6 && decoder.stream.peek_at(0, 6)? == b"Exif\x00\x00" {
trace!("Exif segment present");
// skip bytes we read above
decoder.stream.skip(6)?;
length -= 6;
let exif_bytes = decoder.stream.peek_at(0, length)?.to_vec();
decoder.info.exif_data = Some(exif_bytes);
} else if length > XMP_NAMESPACE_PREFIX.len()
&& decoder.stream.peek_at(0, XMP_NAMESPACE_PREFIX.len())? == XMP_NAMESPACE_PREFIX
{
trace!("XMP Data Present");
decoder.stream.skip(XMP_NAMESPACE_PREFIX.len())?;
length -= XMP_NAMESPACE_PREFIX.len();
let xmp_data = decoder.stream.peek_at(0, length)?.to_vec();
decoder.info.xmp_data = Some(xmp_data);
} else {
warn!("Unknown format for APP1 tag, skipping");
}
decoder.stream.skip(length)?;
Ok(())
}
pub(crate) fn parse_app2<T: ZByteReaderTrait>(
decoder: &mut JpegDecoder<T>
) -> Result<(), DecodeErrors> {
static HDR_META: &[u8] = b"urn:iso:std:iso:ts:21496:-1\0";
static MPF_DATA: &[u8] = b"MPF\0";
let mut length = usize::from(decoder.stream.get_u16_be());
if length < 2 {
return Err(DecodeErrors::FormatStatic("Too small app2 segment"));
}
// length bytes
length -= 2;
if length > 14 && decoder.stream.peek_at(0, 12)? == *b"ICC_PROFILE\0" {
trace!("ICC Profile present");
// skip 12 bytes which indicate ICC profile
length -= 12;
decoder.stream.skip(12)?;
let seq_no = decoder.stream.read_u8();
let num_markers = decoder.stream.read_u8();
// deduct the two bytes we read above
length -= 2;
let data = decoder.stream.peek_at(0, length)?.to_vec();
let icc_chunk = ICCChunk {
seq_no,
num_markers,
data
};
decoder.icc_data.push(icc_chunk);
} else if length > HDR_META.len() && decoder.stream.peek_at(0, HDR_META.len())? == HDR_META {
length = length.saturating_sub(HDR_META.len());
decoder.stream.skip(HDR_META.len())?;
trace!("Gain Map metadata found");
match length {
4 => {
// If gain map metadata length == 4 then here it variables
// https://github.com/google/libultrahdr/blob/bf2aa439eea9ad5da483003fa44182f990f74091/lib/src/jpegr.cpp#L1076C1-L1077C35
// 2 bytes minimum_version: (00 00)
// 2 bytes writer_version: (00 00)
// Perhaps nothing to do with it ?
let _ = decoder.stream.get_u16_be();
let _ = decoder.stream.get_u16_be();
length -= 4;
decoder
.info
.gain_map_info
.push(GainMapInfo { data: Vec::new() });
}
n if n > 4 => {
// If there is perhaps useful gain map info
// we'll read this until end
// https://github.com/google/libultrahdr/blob/bf2aa439eea9ad5da483003fa44182f990f74091/lib/src/jpegr.cpp#L1323
let data = decoder.stream.peek_at(0, length)?.to_vec();
length -= data.len();
decoder.stream.skip(data.len())?;
decoder.info.gain_map_info.push(GainMapInfo { data });
}
_ => {}
}
} else if length > MPF_DATA.len() && decoder.stream.peek_at(0, MPF_DATA.len())? == MPF_DATA {
trace!("MPF Signature present");
length = length.saturating_sub(MPF_DATA.len());
decoder.stream.skip(MPF_DATA.len())?;
// MPF signature taken from here
// https://github.com/google/libultrahdr/blob/bf2aa439eea9ad5da483003fa44182f990f74091/lib/include/ultrahdr/multipictureformat.h#L50
// https://github.com/google/libultrahdr/blob/bf2aa439eea9ad5da483003fa44182f990f74091/lib/src/multipictureformat.cpp#L36
// More info https://www.cipa.jp/std/documents/e/DC-X007-KEY_E.pdf
let data = decoder.stream.peek_at(0, length)?.to_vec();
length -= data.len();
decoder.stream.skip(data.len())?;
decoder.info.multi_picture_information = Some(data);
}
decoder.stream.skip(length)?;
Ok(())
}
/// Small utility function to print Un-zig-zagged quantization tables
fn un_zig_zag<T>(a: &[T]) -> [i32; 64]
where
T: Default + Copy,
i32: core::convert::From<T>
{
let mut output = [i32::default(); 64];
for i in 0..64 {
output[UN_ZIGZAG[i]] = i32::from(a[i]);
}
output
}
+254
View File
@@ -0,0 +1,254 @@
/*
* Copyright (c) 2023.
*
* This software is free software;
*
* You can redistribute it or modify it under terms of the MIT, Apache License or Zlib license
*/
//! This file contains a single struct `HuffmanTable` that
//! stores Huffman tables needed during `BitStream` decoding.
#![allow(clippy::similar_names, clippy::module_name_repetitions)]
use alloc::string::ToString;
use crate::errors::DecodeErrors;
/// Determines how many bits of lookahead we have for our bitstream decoder.
pub const HUFF_LOOKAHEAD: u8 = 9;
/// A struct which contains necessary tables for decoding a JPEG
/// huffman encoded bitstream
pub struct HuffmanTable {
// element `[0]` of each array is unused
/// largest code of length k
pub(crate) maxcode: [i32; 18],
/// offset for codes of length k
/// Answers the question, where do code-lengths of length k end
/// Element 0 is unused
pub(crate) offset: [i32; 18],
/// lookup table for fast decoding
///
/// top bits above HUFF_LOOKAHEAD contain the code length.
///
/// Lower (8) bits contain the symbol in order of increasing code length.
pub(crate) lookup: [i32; 1 << HUFF_LOOKAHEAD],
/// A table which can be used to decode small AC coefficients and
/// do an equivalent of receive_extend
pub(crate) ac_lookup: Option<[i16; 1 << HUFF_LOOKAHEAD]>,
/// Directly represent contents of a JPEG DHT marker
///
/// \# number of symbols with codes of length `k` bits
// bits[0] is unused
/// Symbols in order of increasing code length
pub(crate) values: [u8; 256]
}
impl HuffmanTable {
pub fn new(
codes: &[u8; 17], values: [u8; 256], is_dc: bool, is_progressive: bool
) -> Result<HuffmanTable, DecodeErrors> {
let too_long_code = (i32::from(HUFF_LOOKAHEAD) + 1) << HUFF_LOOKAHEAD;
let mut p = HuffmanTable {
maxcode: [0; 18],
offset: [0; 18],
lookup: [too_long_code; 1 << HUFF_LOOKAHEAD],
values,
ac_lookup: None
};
p.make_derived_table(is_dc, is_progressive, codes)?;
Ok(p)
}
/// Create a new huffman tables with values that aren't fixed
/// used by fill_mjpeg_tables
pub fn new_unfilled(
codes: &[u8; 17], values: &[u8], is_dc: bool, is_progressive: bool
) -> Result<HuffmanTable, DecodeErrors> {
let mut buf = [0; 256];
buf[..values.len()].copy_from_slice(values);
HuffmanTable::new(codes, buf, is_dc, is_progressive)
}
/// Compute derived values for a Huffman table
///
/// This routine performs some validation checks on the table
#[allow(
clippy::cast_possible_truncation,
clippy::cast_possible_wrap,
clippy::cast_sign_loss,
clippy::too_many_lines,
clippy::needless_range_loop
)]
fn make_derived_table(
&mut self, is_dc: bool, _is_progressive: bool, bits: &[u8; 17]
) -> Result<(), DecodeErrors> {
// build a list of code size
let mut huff_size = [0; 257];
// Huffman code lengths
let mut huff_code: [u32; 257] = [0; 257];
// figure C.1 make table of Huffman code length for each symbol
let mut p = 0;
for l in 1..=16 {
let mut i = i32::from(bits[l]);
// table overrun is checked before ,so we dont need to check
while i != 0 {
huff_size[p] = l as u8;
p += 1;
i -= 1;
}
}
huff_size[p] = 0;
let num_symbols = p;
// Generate the codes themselves
// We also validate that the counts represent a legal Huffman code tree
let mut code = 0;
let mut si = i32::from(huff_size[0]);
p = 0;
while huff_size[p] != 0 {
while i32::from(huff_size[p]) == si {
huff_code[p] = code;
code += 1;
p += 1;
}
// maximum code of length si, pre-shifted by 16-k bits
self.maxcode[si as usize] = (code << (16 - si)) as i32;
// code is now 1 more than the last code used for code-length si; but
// it must still fit in si bits, since no code is allowed to be all ones.
if (code as i32) >= (1 << si) {
return Err(DecodeErrors::HuffmanDecode("Bad Huffman Table".to_string()));
}
code <<= 1;
si += 1;
}
// Figure F.15 generate decoding tables for bit-sequential decoding
p = 0;
for l in 0..=16 {
if bits[l] == 0 {
// -1 if no codes of this length
self.maxcode[l] = -1;
} else {
// offset[l]=codes[index of 1st symbol of code length l
// minus minimum code of length l]
self.offset[l] = (p as i32) - (huff_code[p]) as i32;
p += usize::from(bits[l]);
}
}
self.offset[17] = 0;
// we ensure that decode terminates
self.maxcode[17] = 0x000F_FFFF;
/*
* Compute lookahead tables to speed up decoding.
* First we set all the table entries to 0(left justified), indicating "too long";
* (Note too long was set during initialization)
* then we iterate through the Huffman codes that are short enough and
* fill in all the entries that correspond to bit sequences starting
* with that code.
*/
p = 0;
for l in 1..=HUFF_LOOKAHEAD {
for _ in 1..=i32::from(bits[usize::from(l)]) {
// l -> Current code length,
// p => Its index in self.code and self.values
// Generate left justified code followed by all possible bit sequences
let mut look_bits = (huff_code[p] as usize) << (HUFF_LOOKAHEAD - l);
for _ in 0..1 << (HUFF_LOOKAHEAD - l) {
self.lookup[look_bits] =
(i32::from(l) << HUFF_LOOKAHEAD) | i32::from(self.values[p]);
look_bits += 1;
}
p += 1;
}
}
// build an ac table that does an equivalent of decode and receive_extend
if !is_dc {
let mut fast = [255; 1 << HUFF_LOOKAHEAD];
// Iterate over number of symbols
for i in 0..num_symbols {
// get code size for an item
let s = huff_size[i];
if s <= HUFF_LOOKAHEAD {
// if it's lower than what we need for our lookup table create the table
let c = (huff_code[i] << (HUFF_LOOKAHEAD - s)) as usize;
let m = (1 << (HUFF_LOOKAHEAD - s)) as usize;
for j in 0..m {
fast[c + j] = i as i16;
}
}
}
// build a table that decodes both magnitude and value of small ACs in
// one go.
let mut fast_ac = [0; 1 << HUFF_LOOKAHEAD];
for i in 0..(1 << HUFF_LOOKAHEAD) {
let fast_v = fast[i];
if fast_v < 255 {
// get symbol value from AC table
let rs = self.values[fast_v as usize];
// shift by 4 to get run length
let run = i16::from((rs >> 4) & 15);
// get magnitude bits stored at the lower 3 bits
let mag_bits = i16::from(rs & 15);
// length of the bit we've read
let len = i16::from(huff_size[fast_v as usize]);
if mag_bits != 0 && (len + mag_bits) <= i16::from(HUFF_LOOKAHEAD) {
// magnitude code followed by receive_extend code
let mut k = (((i as i16) << len) & ((1 << HUFF_LOOKAHEAD) - 1))
>> (i16::from(HUFF_LOOKAHEAD) - mag_bits);
let m = 1 << (mag_bits - 1);
if k < m {
k += (!0_i16 << mag_bits) + 1;
};
// if result is small enough fit into fast ac table
if (-128..=127).contains(&k) {
fast_ac[i] = (k << 8) + (run << 4) + (len + mag_bits);
}
}
}
}
self.ac_lookup = Some(fast_ac);
}
// Validate symbols as being reasonable
// For AC tables, we make no check, but accept all byte values 0..255
// For DC tables, we require symbols to be in range 0..15
if is_dc {
for i in 0..num_symbols {
let sym = self.values[i];
if sym > 15 {
return Err(DecodeErrors::HuffmanDecode("Bad Huffman Table".to_string()));
}
}
}
Ok(())
}
}
+206
View File
@@ -0,0 +1,206 @@
/*
* Copyright (c) 2023.
*
* This software is free software;
*
* You can redistribute it or modify it under terms of the MIT, Apache License or Zlib license
*/
//! Routines for IDCT
//!
//! Essentially we provide 2 routines for IDCT, a scalar implementation and a not super optimized
//! AVX2 one, i'll talk about them here.
//!
//! There are 2 reasons why we have the avx one
//! 1. No one compiles with -C target-features=avx2 hence binaries won't probably take advantage(even
//! if it exists).
//! 2. AVX employs zero short circuit in a way the scalar code cannot employ it.
//! - AVX does this by checking for MCU's whose 63 AC coefficients are zero and if true, it writes
//! values directly, if false, it goes the long way of calculating.
//! - Although this can be trivially implemented in the scalar version, it generates code
//! I'm not happy width(scalar version that basically loops and that is too many branches for me)
//! The avx one does a better job of using bitwise or's with (`_mm256_or_si256`) which is magnitudes of faster
//! than anything I could come up with
//!
//! The AVX code also has some cool transpose_u16 instructions which look so complicated to be cool
//! (spoiler alert, i barely understand how it works, that's why I credited the owner).
//!
#![allow(
clippy::excessive_precision,
clippy::unreadable_literal,
clippy::module_name_repetitions,
unused_parens,
clippy::wildcard_imports
)]
use zune_core::log::debug;
use zune_core::options::DecoderOptions;
use crate::decoder::IDCTPtr;
use crate::idct::scalar::{idct_int, idct_int_1x1};
#[cfg(feature = "x86")]
pub mod avx2;
#[cfg(feature = "neon")]
pub mod neon;
pub mod scalar;
/// Choose an appropriate IDCT function
#[allow(unused_variables)]
pub fn choose_idct_func(options: &DecoderOptions) -> IDCTPtr {
#[cfg(any(target_arch = "x86", target_arch = "x86_64"))]
#[cfg(feature = "x86")]
{
if options.use_avx2() {
debug!("Using vector integer IDCT");
return |a: &mut [i32; 64], b: &mut [i16], c: usize| {
// SAFETY: `options.use_avx2()` only returns true if avx2 is supported.
unsafe { avx2::idct_avx2(a,b,c) }
};
}
}
#[cfg(target_arch = "aarch64")]
#[cfg(feature = "neon")]
{
if options.use_neon() {
debug!("Using vector integer IDCT");
return |a: &mut [i32; 64], b: &mut [i16], c: usize| {
// SAFETY: `options.use_neon()` only returns true if neon is supported.
unsafe { neon::idct_neon(a,b,c) }
};
}
}
debug!("Using scalar integer IDCT");
// use generic one
return idct_int;
}
/// Choose a function to implement 4x4 IDCT.
///
/// These functions get the same input but have an extra contract: Only the first 4x4 block of
/// coefficients are non-zero. All other entries are zeroed.
///
/// **The callee must uphold that contract on return**
pub fn choose_idct_4x4_func(_options: &DecoderOptions) -> IDCTPtr {
#[cfg(any(target_arch = "x86", target_arch = "x86_64"))]
#[cfg(feature = "x86")]
{
if _options.use_avx2() {
debug!("Using vector integer IDCT");
return |a: &mut [i32; 64], b: &mut [i16], c: usize| {
// SAFETY: `options.use_avx2()` only returns true if avx2 is supported.
unsafe { avx2::idct_avx2_4x4(a,b,c) }
};
}
}
scalar::idct4x4
}
pub fn choose_idct_1x1_func(_: &DecoderOptions) -> IDCTPtr {
// These are simple stores, no alternative implementation for now
idct_int_1x1
}
#[cfg(test)]
#[allow(unreachable_code)]
#[allow(dead_code)]
mod tests {
use super::*;
#[test]
fn idct_test0() {
let stride = 8;
let mut coeff = [10; 64];
let mut coeff2 = [10; 64];
let mut output_scalar = [0; 64];
let mut output_vector = [0; 64];
let idct_func = choose_idct_func(&DecoderOptions::new_fast());
idct_func(&mut coeff, &mut output_vector, stride);
idct_int(&mut coeff2, &mut output_scalar, stride);
assert_eq!(output_scalar, output_vector, "IDCT and scalar do not match");
}
#[test]
fn do_idct_test1() {
let stride = 8;
let mut coeff = [14; 64];
let mut coeff2 = [14; 64];
let mut output_scalar = [0; 64];
let mut output_vector = [0; 64];
let idct_func = choose_idct_func(&DecoderOptions::new_fast());
idct_func(&mut coeff, &mut output_vector, stride);
idct_int(&mut coeff2, &mut output_scalar, stride);
assert_eq!(output_scalar, output_vector, "IDCT and scalar do not match");
}
#[test]
fn do_idct_test2() {
let stride = 8;
let mut coeff = [0; 64];
coeff[0] = 255;
coeff[63] = -256;
let mut coeff2 = coeff;
let mut output_scalar = [0; 64];
let mut output_vector = [0; 64];
let idct_func = choose_idct_func(&DecoderOptions::new_fast());
idct_func(&mut coeff, &mut output_vector, stride);
idct_int(&mut coeff2, &mut output_scalar, stride);
assert_eq!(output_scalar, output_vector, "IDCT and scalar do not match");
}
#[test]
fn do_idct_zeros() {
let stride = 8;
let mut coeff = [0; 64];
let mut coeff2 = [0; 64];
let mut output_scalar = [0; 64];
let mut output_vector = [0; 64];
let idct_func = choose_idct_func(&DecoderOptions::new_fast());
idct_func(&mut coeff, &mut output_vector, stride);
idct_int(&mut coeff2, &mut output_scalar, stride);
assert_eq!(output_scalar, output_vector, "IDCT and scalar do not match");
}
#[test]
fn idct_4x4() {
#[rustfmt::skip]
const A: [i32; 32] = [
-254, -7, 0, 0, 0, 0, 0, 0,
7, 0, -30, 32, 0, 0, 0, 0,
7, 0, -30, 32, 0, 0, 0, 0,
7, 0, -30, 32, 0, 0, 0, 0,
];
let v: Vec<IDCTPtr> = vec![
choose_idct_func(&DecoderOptions::new_safe()),
choose_idct_4x4_func(&DecoderOptions::new_safe()),
choose_idct_func(&DecoderOptions::new_fast()),
choose_idct_4x4_func(&DecoderOptions::new_fast()),
];
let dct_names = vec![
"safe idct",
"safe idct 4x4",
"fast idct",
"fast idct 4x4",
];
let mut color = vec![];
for idct in v {
let mut a = [0i32; 64];
a[..32].copy_from_slice(&A);
let mut b = [0i16; 64];
idct(&mut a, &mut b, 8);
color.push(b);
}
for (wnd, name) in color.windows(2).zip(&dct_names) {
let [a, b] = wnd else { unreachable!() };
assert_eq!(a, b, "{name}");
}
}
}
+398
View File
@@ -0,0 +1,398 @@
/*
* Copyright (c) 2023.
*
* This software is free software;
*
* You can redistribute it or modify it under terms of the MIT, Apache License or Zlib license
*/
#![cfg(any(target_arch = "x86", target_arch = "x86_64"))]
//! AVX optimised IDCT.
//!
//! Okay not thaat optimised.
//!
//!
//! # The implementation
//! The implementation is neatly broken down into two operations.
//!
//! 1. Test for zeroes
//! > There is a shortcut method for idct where when all AC values are zero, we can get the answer really quickly.
//! by scaling the 1/8th of the DCT coefficient of the block to the whole block and level shifting.
//!
//! 2. If above fails, we proceed to carry out IDCT as a two pass one dimensional algorithm.
//! IT does two whole scans where it carries out IDCT on all items
//! After each successive scan, data is transposed in register(thank you x86 SIMD powers). and the second
//! pass is carried out.
//!
//! The code is not super optimized, it produces bit identical results with scalar code hence it's
//! `mm256_add_epi16`
//! and it also has the advantage of making this implementation easy to maintain.
#![cfg(feature = "x86")]
#![allow(dead_code)]
#[cfg(target_arch = "x86")]
use core::arch::x86::*;
#[cfg(target_arch = "x86_64")]
use core::arch::x86_64::*;
use crate::unsafe_utils::{transpose, YmmRegister};
const SCALE_BITS: i32 = 512 + 65536 + (128 << 17);
// Pack i32 to i16's,
// clamp them to be between 0-255
// Undo shuffling
// Store back to array
macro_rules! permute_store {
($x:tt,$y:tt,$index:tt,$out:tt,$stride:tt) => {
let a = _mm256_packs_epi32($x, $y);
// Clamp the values after packing, we can clamp more values at once
let b = clamp_avx(a);
// /Undo shuffling
let c = _mm256_permute4x64_epi64(b, shuffle(3, 1, 2, 0));
// store first vector
_mm_storeu_si128(
($out)
.get_mut($index..$index + 8)
.unwrap()
.as_mut_ptr()
.cast(),
_mm256_extractf128_si256::<0>(c),
);
$index += $stride;
// second vector
_mm_storeu_si128(
($out)
.get_mut($index..$index + 8)
.unwrap()
.as_mut_ptr()
.cast(),
_mm256_extractf128_si256::<1>(c),
);
$index += $stride;
};
}
#[target_feature(enable = "avx2")]
#[allow(
clippy::too_many_lines,
clippy::cast_possible_truncation,
clippy::similar_names,
clippy::op_ref,
unused_assignments,
clippy::zero_prefixed_literal
)]
pub unsafe fn idct_avx2(
in_vector: &mut [i32; 64], out_vector: &mut [i16], stride: usize,
) {
let mut pos = 0;
// load into registers
//
// We sign extend i16's to i32's and calculate them with extended precision and
// later reduce them to i16's when we are done carrying out IDCT
let rw0 = _mm256_loadu_si256(in_vector[00..].as_ptr().cast());
let rw1 = _mm256_loadu_si256(in_vector[08..].as_ptr().cast());
let rw2 = _mm256_loadu_si256(in_vector[16..].as_ptr().cast());
let rw3 = _mm256_loadu_si256(in_vector[24..].as_ptr().cast());
let rw4 = _mm256_loadu_si256(in_vector[32..].as_ptr().cast());
let rw5 = _mm256_loadu_si256(in_vector[40..].as_ptr().cast());
let rw6 = _mm256_loadu_si256(in_vector[48..].as_ptr().cast());
let rw7 = _mm256_loadu_si256(in_vector[56..].as_ptr().cast());
// Forward DCT and quantization may cause all the AC terms to be zero, for such
// cases we can try to accelerate it
// Basically the poop is that whenever the array has 63 zeroes, its idct is
// (arr[0]>>3)or (arr[0]/8) propagated to all the elements.
// We first test to see if the array contains zero elements and if it does, we go the
// short way.
//
// This reduces IDCT overhead from about 39% to 18 %, almost half
// Do another load for the first row, we don't want to check DC value, because
// we only care about AC terms
let rw8 = _mm256_loadu_si256(in_vector[1..].as_ptr().cast());
let mut bitmap = _mm256_or_si256(rw1, rw2);
bitmap = _mm256_or_si256(bitmap, rw3);
bitmap = _mm256_or_si256(bitmap, rw4);
bitmap = _mm256_or_si256(bitmap, rw5);
bitmap = _mm256_or_si256(bitmap, rw6);
bitmap = _mm256_or_si256(bitmap, rw7);
bitmap = _mm256_or_si256(bitmap, rw8);
if _mm256_testz_si256(bitmap, bitmap) == 1 {
// AC terms all zero, idct of the block is ( coeff[0] * qt[0] )/8 + 128 (bias)
// (and clamped to 255)
// Round by adding 0.5 * (1 << 3) and offset by adding (128 << 3) before scaling
let coeff = ((in_vector[0] + 4 + 1024) >> 3).clamp(0, 255) as i16;
let idct_value = _mm_set1_epi16(coeff);
macro_rules! store {
($pos:tt,$value:tt) => {
// store
_mm_storeu_si128(
out_vector
.get_mut($pos..$pos + 8)
.unwrap()
.as_mut_ptr()
.cast(),
$value,
);
$pos += stride;
};
}
store!(pos, idct_value);
store!(pos, idct_value);
store!(pos, idct_value);
store!(pos, idct_value);
store!(pos, idct_value);
store!(pos, idct_value);
store!(pos, idct_value);
store!(pos, idct_value);
return;
}
let mut row0 = YmmRegister { mm256: rw0 };
let mut row1 = YmmRegister { mm256: rw1 };
let mut row2 = YmmRegister { mm256: rw2 };
let mut row3 = YmmRegister { mm256: rw3 };
let mut row4 = YmmRegister { mm256: rw4 };
let mut row5 = YmmRegister { mm256: rw5 };
let mut row6 = YmmRegister { mm256: rw6 };
let mut row7 = YmmRegister { mm256: rw7 };
macro_rules! dct_pass {
($SCALE_BITS:tt,$scale:tt) => {
// There are a lot of ways to do this
// but to keep it simple(and beautiful), ill make a direct translation of the
// scalar code to also make this code fully transparent(this version and the non
// avx one should produce identical code.)
// even part
let p1 = (row2 + row6) * 2217;
let mut t2 = p1 + row6 * -7567;
let mut t3 = p1 + row2 * 3135;
let mut t0 = YmmRegister {
mm256: _mm256_slli_epi32((row0 + row4).mm256, 12),
};
let mut t1 = YmmRegister {
mm256: _mm256_slli_epi32((row0 - row4).mm256, 12),
};
let x0 = t0 + t3 + $SCALE_BITS;
let x3 = t0 - t3 + $SCALE_BITS;
let x1 = t1 + t2 + $SCALE_BITS;
let x2 = t1 - t2 + $SCALE_BITS;
let p3 = row7 + row3;
let p4 = row5 + row1;
let p1 = row7 + row1;
let p2 = row5 + row3;
let p5 = (p3 + p4) * 4816;
t0 = row7 * 1223;
t1 = row5 * 8410;
t2 = row3 * 12586;
t3 = row1 * 6149;
let p1 = p5 + p1 * -3685;
let p2 = p5 + (p2 * -10497);
let p3 = p3 * -8034;
let p4 = p4 * -1597;
t3 += p1 + p4;
t2 += p2 + p3;
t1 += p2 + p4;
t0 += p1 + p3;
row0.mm256 = _mm256_srai_epi32((x0 + t3).mm256, $scale);
row1.mm256 = _mm256_srai_epi32((x1 + t2).mm256, $scale);
row2.mm256 = _mm256_srai_epi32((x2 + t1).mm256, $scale);
row3.mm256 = _mm256_srai_epi32((x3 + t0).mm256, $scale);
row4.mm256 = _mm256_srai_epi32((x3 - t0).mm256, $scale);
row5.mm256 = _mm256_srai_epi32((x2 - t1).mm256, $scale);
row6.mm256 = _mm256_srai_epi32((x1 - t2).mm256, $scale);
row7.mm256 = _mm256_srai_epi32((x0 - t3).mm256, $scale);
};
}
// Process rows
dct_pass!(512, 10);
transpose(
&mut row0, &mut row1, &mut row2, &mut row3, &mut row4, &mut row5, &mut row6, &mut row7,
);
// process columns
dct_pass!(SCALE_BITS, 17);
transpose(
&mut row0, &mut row1, &mut row2, &mut row3, &mut row4, &mut row5, &mut row6, &mut row7,
);
// Pack and write the values back to the array
permute_store!((row0.mm256), (row1.mm256), pos, out_vector, stride);
permute_store!((row2.mm256), (row3.mm256), pos, out_vector, stride);
permute_store!((row4.mm256), (row5.mm256), pos, out_vector, stride);
permute_store!((row6.mm256), (row7.mm256), pos, out_vector, stride);
}
#[target_feature(enable = "avx2")]
#[allow(
clippy::too_many_lines,
clippy::cast_possible_truncation,
clippy::similar_names,
clippy::op_ref,
unused_assignments,
clippy::zero_prefixed_literal
)]
pub unsafe fn idct_avx2_4x4(
in_vector: &mut [i32; 64], out_vector: &mut [i16], stride: usize,
) {
let rw0 = _mm256_loadu_si256(in_vector[00..].as_ptr().cast());
let rw1 = _mm256_loadu_si256(in_vector[08..].as_ptr().cast());
let rw2 = _mm256_loadu_si256(in_vector[16..].as_ptr().cast());
let rw3 = _mm256_loadu_si256(in_vector[24..].as_ptr().cast());
let mut row0 = YmmRegister { mm256: rw0 };
let mut row1 = YmmRegister { mm256: rw1 };
let mut row2 = YmmRegister { mm256: rw2 };
let mut row3 = YmmRegister { mm256: rw3 };
let mut row4 = YmmRegister { mm256: rw0 };
let mut row5 = YmmRegister { mm256: rw0 };
let mut row6 = YmmRegister { mm256: rw0 };
let mut row7 = YmmRegister { mm256: rw0 };
{
row0.mm256 = _mm256_slli_epi32(row0.mm256, 12);
row0 += 512;
let i2 = row2;
let p1 = i2 * 2217;
let p3 = i2 * 5352;
let x0 = row0 + p3;
let x1 = row0 + p1;
let x2 = row0 - p1;
let x3 = row0 - p3;
// odd part
let i4 = row3;
let i3 = row1;
let p5 = (i4 + i3) * 4816;
let p1 = p5 + i3 * -3685;
let p2 = p5 + i4 * -10497;
let t3 = p5 + i3 * 867;
let t2 = p5 + i4 * -5945;
let t1 = p2 + i3 * -1597;
let t0 = p1 + i4 * -8034;
row0.mm256 = _mm256_srai_epi32((x0 + t3).mm256, 10);
row1.mm256 = _mm256_srai_epi32((x1 + t2).mm256, 10);
row2.mm256 = _mm256_srai_epi32((x2 + t1).mm256, 10);
row3.mm256 = _mm256_srai_epi32((x3 + t0).mm256, 10);
row4.mm256 = _mm256_srai_epi32((x3 - t0).mm256, 10);
row5.mm256 = _mm256_srai_epi32((x2 - t1).mm256, 10);
row6.mm256 = _mm256_srai_epi32((x1 - t2).mm256, 10);
row7.mm256 = _mm256_srai_epi32((x0 - t3).mm256, 10);
}
transpose(
&mut row0, &mut row1, &mut row2, &mut row3, &mut row4, &mut row5, &mut row6, &mut row7,
);
{
let i2 = row2;
let i0 = row0;
row0.mm256 = _mm256_slli_epi32(i0.mm256, 12);
let t0 = row0 + SCALE_BITS;
let t2 = i2 * 2217;
let t3 = i2 * 5352;
// constants scaled things up by 1<<12, plus we had 1<<2 from first
// loop, plus horizontal and vertical each scale by sqrt(8) so together
// we've got an extra 1<<3, so 1<<17 total we need to remove.
// so we want to round that, which means adding 0.5 * 1<<17,
// aka 65536. Also, we'll end up with -128 to 127 that we want
// to encode as 0..255 by adding 128, so we'll add that before the shift
// Rounding constant is already added into `t0`
let x0 = t0 + t3;
let x3 = t0 - t3;
let x1 = t0 + t2;
let x2 = t0 - t2;
// odd part
let i3 = row3;
let i1 = row1;
let p5 = (i3 + i1) * 4816;
let p1 = p5 + i1 * -3685;
let p2 = p5 + i3 * -10497;
let t3 = p5 + i1 * 867;
let t2 = p5 + i3 * -5945;
let t1 = p2 + i1 * -1597;
let t0 = p1 + i3 * -8034;
row0.mm256 = _mm256_srai_epi32((x0 + t3).mm256, 17);
row1.mm256 = _mm256_srai_epi32((x1 + t2).mm256, 17);
row2.mm256 = _mm256_srai_epi32((x2 + t1).mm256, 17);
row3.mm256 = _mm256_srai_epi32((x3 + t0).mm256, 17);
row4.mm256 = _mm256_srai_epi32((x3 - t0).mm256, 17);
row5.mm256 = _mm256_srai_epi32((x2 - t1).mm256, 17);
row6.mm256 = _mm256_srai_epi32((x1 - t2).mm256, 17);
row7.mm256 = _mm256_srai_epi32((x0 - t3).mm256, 17);
}
transpose(
&mut row0, &mut row1, &mut row2, &mut row3, &mut row4, &mut row5, &mut row6, &mut row7,
);
let mut pos = 0;
// Pack and write the values back to the array
permute_store!((row0.mm256), (row1.mm256), pos, out_vector, stride);
permute_store!((row2.mm256), (row3.mm256), pos, out_vector, stride);
permute_store!((row4.mm256), (row5.mm256), pos, out_vector, stride);
permute_store!((row6.mm256), (row7.mm256), pos, out_vector, stride);
}
#[inline]
#[target_feature(enable = "avx2")]
unsafe fn clamp_avx(reg: __m256i) -> __m256i {
let min_s = _mm256_set1_epi16(0);
let max_s = _mm256_set1_epi16(255);
let max_v = _mm256_max_epi16(reg, min_s); //max(a,0)
let min_v = _mm256_min_epi16(max_v, max_s); //min(max(a,0),255)
return min_v;
}
/// A copy of `_MM_SHUFFLE()` that doesn't require
/// a nightly compiler
#[inline]
const fn shuffle(z: i32, y: i32, x: i32, w: i32) -> i32 {
((z << 6) | (y << 4) | (x << 2) | w)
}
+280
View File
@@ -0,0 +1,280 @@
/*
* Copyright (c) 2023.
*
* This software is free software;
*
* You can redistribute it or modify it under terms of the MIT, Apache License or Zlib license
*/
#![cfg(target_arch = "aarch64")]
//! AVX optimised IDCT.
//!
//! Okay not thaat optimised.
//!
//!
//! # The implementation
//! The implementation is neatly broken down into two operations.
//!
//! 1. Test for zeroes
//! > There is a shortcut method for idct where when all AC values are zero, we can get the answer really quickly.
//! by scaling the 1/8th of the DCT coefficient of the block to the whole block and level shifting.
//!
//! 2. If above fails, we proceed to carry out IDCT as a two pass one dimensional algorithm.
//! IT does two whole scans where it carries out IDCT on all items
//! After each successive scan, data is transposed in register(thank you x86 SIMD powers). and the second
//! pass is carried out.
//!
//! The code is not super optimized, it produces bit identical results with scalar code hence it's
//! `mm256_add_epi16`
//! and it also has the advantage of making this implementation easy to maintain.
#![cfg(feature = "neon")]
use core::arch::aarch64::*;
use crate::unsafe_utils::{transpose, YmmRegister};
const SCALE_BITS: i32 = 512 + 65536 + (128 << 17);
#[inline]
#[target_feature(enable = "neon")]
unsafe fn pack_16(a: int32x4x2_t) -> int16x8_t {
vcombine_s16(vqmovn_s32(a.0), vqmovn_s32(a.1))
}
#[inline]
#[target_feature(enable = "neon")]
unsafe fn condense_bottom_16(a: int32x4x2_t, b: int32x4x2_t) -> int16x8x2_t {
int16x8x2_t(pack_16(a), pack_16(b))
}
#[target_feature(enable = "neon")]
#[allow(
clippy::too_many_lines,
clippy::cast_possible_truncation,
clippy::similar_names,
clippy::op_ref,
unused_assignments,
clippy::zero_prefixed_literal
)]
pub unsafe fn idct_neon(
in_vector: &mut [i32; 64], out_vector: &mut [i16], stride: usize
) {
let mut pos = 0;
// load into registers
//
// We sign extend i16's to i32's and calculate them with extended precision and
// later reduce them to i16's when we are done carrying out IDCT
let mut row0 = YmmRegister::load(in_vector[00..].as_ptr().cast());
let mut row1 = YmmRegister::load(in_vector[08..].as_ptr().cast());
let mut row2 = YmmRegister::load(in_vector[16..].as_ptr().cast());
let mut row3 = YmmRegister::load(in_vector[24..].as_ptr().cast());
let mut row4 = YmmRegister::load(in_vector[32..].as_ptr().cast());
let mut row5 = YmmRegister::load(in_vector[40..].as_ptr().cast());
let mut row6 = YmmRegister::load(in_vector[48..].as_ptr().cast());
let mut row7 = YmmRegister::load(in_vector[56..].as_ptr().cast());
// Forward DCT and quantization may cause all the AC terms to be zero, for such
// cases we can try to accelerate it
// Basically the poop is that whenever the array has 63 zeroes, its idct is
// (arr[0]>>3)or (arr[0]/8) propagated to all the elements.
// We first test to see if the array contains zero elements and if it does, we go the
// short way.
//
// This reduces IDCT overhead from about 39% to 18 %, almost half
// Do another load for the first row, we don't want to check DC value, because
// we only care about AC terms
// TODO this should be a shift/shuffle, not a likely unaligned load
let row8 = YmmRegister::load(in_vector[1..].as_ptr().cast());
let or_tree = (((row1 | row8) | (row2 | row3)) | ((row4 | row5) | (row6 | row7)));
if or_tree.all_zero() {
// AC terms all zero, idct of the block is ( coeff[0] * qt[0] )/8 + 128 (bias)
// (and clamped to 255)
// Round by adding 0.5 * (1 << 3) and offset by adding (128 << 3) before scaling
let coeff = ((in_vector[0] + 4 + 1024) >> 3).clamp(0, 255) as i16;
let idct_value = vdupq_n_s16(coeff);
macro_rules! store {
($pos:tt,$value:tt) => {
// store
vst1q_s16(
out_vector
.get_mut($pos..$pos + 8)
.unwrap()
.as_mut_ptr()
.cast(),
$value
);
$pos += stride;
};
}
store!(pos, idct_value);
store!(pos, idct_value);
store!(pos, idct_value);
store!(pos, idct_value);
store!(pos, idct_value);
store!(pos, idct_value);
store!(pos, idct_value);
store!(pos, idct_value);
return;
}
macro_rules! dct_pass {
($SCALE_BITS:tt,$scale:tt) => {
// There are a lot of ways to do this
// but to keep it simple(and beautiful), ill make a direct translation of the
// scalar code to also make this code fully transparent(this version and the non
// avx one should produce identical code.)
// Compiler does a pretty good job of optimizing add + mul pairs
// into multiply-acumulate pairs
// even part
let p1 = (row2 + row6) * 2217;
let mut t2 = p1 + row6 * -7567;
let mut t3 = p1 + row2 * 3135;
let mut t0 = (row0 + row4).const_shl::<12>();
let mut t1 = (row0 - row4).const_shl::<12>();
let x0 = t0 + t3 + $SCALE_BITS;
let x3 = t0 - t3 + $SCALE_BITS;
let x1 = t1 + t2 + $SCALE_BITS;
let x2 = t1 - t2 + $SCALE_BITS;
let p3 = row7 + row3;
let p4 = row5 + row1;
let p1 = row7 + row1;
let p2 = row5 + row3;
let p5 = (p3 + p4) * 4816;
t0 = row7 * 1223;
t1 = row5 * 8410;
t2 = row3 * 12586;
t3 = row1 * 6149;
let p1 = p5 + p1 * -3685;
let p2 = p5 + (p2 * -10497);
let p3 = p3 * -8034;
let p4 = p4 * -1597;
t3 += p1 + p4;
t2 += p2 + p3;
t1 += p2 + p4;
t0 += p1 + p3;
row0 = (x0 + t3).const_shra::<$scale>();
row1 = (x1 + t2).const_shra::<$scale>();
row2 = (x2 + t1).const_shra::<$scale>();
row3 = (x3 + t0).const_shra::<$scale>();
row4 = (x3 - t0).const_shra::<$scale>();
row5 = (x2 - t1).const_shra::<$scale>();
row6 = (x1 - t2).const_shra::<$scale>();
row7 = (x0 - t3).const_shra::<$scale>();
};
}
// Process rows
dct_pass!(512, 10);
transpose(
&mut row0, &mut row1, &mut row2, &mut row3, &mut row4, &mut row5, &mut row6, &mut row7
);
// process columns
dct_pass!(SCALE_BITS, 17);
transpose(
&mut row0, &mut row1, &mut row2, &mut row3, &mut row4, &mut row5, &mut row6, &mut row7
);
// Pack i32 to i16's,
// clamp them to be between 0-255
// Undo shuffling
// Store back to array
// This could potentially be reorganized to take advantage of the multi-register stores
macro_rules! permute_store {
($x:tt,$y:tt,$index:tt,$out:tt) => {
let a = condense_bottom_16($x, $y);
// Clamp the values after packing, we can clamp more values at once
let b = clamp256_neon(a);
// store first vector
vst1q_s16(
($out)
.get_mut($index..$index + 8)
.unwrap()
.as_mut_ptr()
.cast(),
b.0
);
$index += stride;
// second vector
vst1q_s16(
($out)
.get_mut($index..$index + 8)
.unwrap()
.as_mut_ptr()
.cast(),
b.1
);
$index += stride;
};
}
// Pack and write the values back to the array
permute_store!((row0.mm256), (row1.mm256), pos, out_vector);
permute_store!((row2.mm256), (row3.mm256), pos, out_vector);
permute_store!((row4.mm256), (row5.mm256), pos, out_vector);
permute_store!((row6.mm256), (row7.mm256), pos, out_vector);
}
#[inline]
#[target_feature(enable = "neon")]
unsafe fn clamp_neon(reg: int16x8_t) -> int16x8_t {
let min_s = vdupq_n_s16(0);
let max_s = vdupq_n_s16(255);
let max_v = vmaxq_s16(reg, min_s); //max(a,0)
let min_v = vminq_s16(max_v, max_s); //min(max(a,0),255)
min_v
}
#[inline]
#[target_feature(enable = "neon")]
unsafe fn clamp256_neon(reg: int16x8x2_t) -> int16x8x2_t {
int16x8x2_t(clamp_neon(reg.0), clamp_neon(reg.1))
}
#[cfg(test)]
mod test {
use super::*;
#[test]
fn test_neon_clamp_256() {
unsafe {
let vals: [i16; 16] = [-1, -2, -3, 4, 256, 257, 258, 240, -1, 290, 2, 3, 4, 5, 6, 7];
let loaded = vld1q_s16_x2(vals.as_ptr().cast());
let shuffled = clamp256_neon(loaded);
let mut result: [i16; 16] = [0; 16];
vst1q_s16_x2(result.as_mut_ptr().cast(), shuffled);
assert_eq!(
result,
[0, 0, 0, 4, 255, 255, 255, 240, 0, 255, 2, 3, 4, 5, 6, 7]
)
}
}
}
+293
View File
@@ -0,0 +1,293 @@
/*
* Copyright (c) 2023.
*
* This software is free software;
*
* You can redistribute it or modify it under terms of the MIT, Apache License or Zlib license
*/
//! Platform independent IDCT algorithm
//!
//! Not as fast as AVX one.
const SCALE_BITS: i32 = 512 + 65536 + (128 << 17);
#[inline(always)]
fn wa(a: i32, b: i32) -> i32 {
a.wrapping_add(b)
}
#[inline(always)]
fn ws(a: i32, b: i32) -> i32 {
a.wrapping_sub(b)
}
#[inline(always)]
fn wm(a: i32, b: i32) -> i32 {
a.wrapping_mul(b)
}
#[inline]
pub fn idct_int_1x1(in_vector: &mut [i32; 64], mut out_vector: &mut [i16], stride: usize) {
let coeff = ((wa(wa(in_vector[0], 4), 1024) >> 3).clamp(0, 255)) as i16;
out_vector[..8].fill(coeff);
for _ in 0..7 {
out_vector = &mut out_vector[stride..];
out_vector[..8].fill(coeff);
}
}
#[allow(unused_assignments)]
#[allow(
clippy::too_many_lines,
clippy::op_ref,
clippy::cast_possible_truncation
)]
pub fn idct_int(in_vector: &mut [i32; 64], out_vector: &mut [i16], stride: usize) {
let mut pos = 0;
let mut i = 0;
if &in_vector[1..] == &[0_i32; 63] {
return idct_int_1x1(in_vector, out_vector, stride);
}
// vertical pass
for ptr in 0..8 {
let p2 = in_vector[ptr + 16];
let p3 = in_vector[ptr + 48];
let p1 = wm(wa(p2, p3), 2217);
let t2 = wa(p1, wm(p3, -7567));
let t3 = wa(p1, wm(p2, 3135));
let p2 = in_vector[ptr];
let p3 = in_vector[32 + ptr];
let t0 = fsh(wa(p2, p3));
let t1 = fsh(ws(p2, p3));
let x0 = wa(wa(t0, t3), 512);
let x3 = wa(ws(t0, t3), 512);
let x1 = wa(wa(t1, t2), 512);
let x2 = wa(ws(t1, t2), 512);
let mut t0 = in_vector[ptr + 56];
let mut t1 = in_vector[ptr + 40];
let mut t2 = in_vector[ptr + 24];
let mut t3 = in_vector[ptr + 8];
let p3 = wa(t0, t2);
let p4 = wa(t1, t3);
let p1 = wa(t0, t3);
let p2 = wa(t1, t2);
let p5 = wm(wa(p3, p4), 4816);
t0 = wm(t0, 1223);
t1 = wm(t1, 8410);
t2 = wm(t2, 12586);
t3 = wm(t3, 6149);
let p1 = wa(p5, wm(p1, -3685));
let p2 = wa(p5, wm(p2, -10497));
let p3 = wm(p3, -8034);
let p4 = wm(p4, -1597);
t3 = wa(t3, wa(p1, p4));
t2 = wa(t2, wa(p2, p3));
t1 = wa(t1, wa(p2, p4));
t0 = wa(t0, wa(p1, p3));
in_vector[ptr] = ws(wa(x0, t3), 0) >> 10;
in_vector[ptr + 8] = ws(wa(x1, t2), 0) >> 10;
in_vector[ptr + 16] = ws(wa(x2, t1), 0) >> 10;
in_vector[ptr + 24] = ws(wa(x3, t0), 0) >> 10;
in_vector[ptr + 32] = ws(ws(x3, t0), 0) >> 10;
in_vector[ptr + 40] = ws(ws(x2, t1), 0) >> 10;
in_vector[ptr + 48] = ws(ws(x1, t2), 0) >> 10;
in_vector[ptr + 56] = ws(ws(x0, t3), 0) >> 10;
}
// horizontal pass
while i < 64 {
let p2 = in_vector[i + 2];
let p3 = in_vector[i + 6];
let p1 = wm(wa(p2, p3), 2217);
let t2 = wa(p1, wm(p3, -7567));
let t3 = wa(p1, wm(p2, 3135));
let p2 = in_vector[i];
let p3 = in_vector[i + 4];
let t0 = fsh(wa(p2, p3));
let t1 = fsh(ws(p2, p3));
let x0 = wa(wa(t0, t3), SCALE_BITS);
let x3 = wa(ws(t0, t3), SCALE_BITS);
let x1 = wa(wa(t1, t2), SCALE_BITS);
let x2 = wa(ws(t1, t2), SCALE_BITS);
let mut t0 = in_vector[i + 7];
let mut t1 = in_vector[i + 5];
let mut t2 = in_vector[i + 3];
let mut t3 = in_vector[i + 1];
let p3 = wa(t0, t2);
let p4 = wa(t1, t3);
let p1 = wa(t0, t3);
let p2 = wa(t1, t2);
let p5 = wm(wa(p3, p4), f2f(1.175875602));
t0 = wm(t0, 1223);
t1 = wm(t1, 8410);
t2 = wm(t2, 12586);
t3 = wm(t3, 6149);
let p1 = wa(p5, wm(p1, -3685));
let p2 = wa(p5, wm(p2, -10497));
let p3 = wm(p3, -8034);
let p4 = wm(p4, -1597);
t3 = wa(t3, wa(p1, p4));
t2 = wa(t2, wa(p2, p3));
t1 = wa(t1, wa(p2, p4));
t0 = wa(t0, wa(p1, p3));
let out: &mut [i16; 8] = out_vector
.get_mut(pos..pos + 8)
.unwrap()
.try_into()
.unwrap();
out[0] = clamp(wa(x0, t3) >> 17);
out[1] = clamp(wa(x1, t2) >> 17);
out[2] = clamp(wa(x2, t1) >> 17);
out[3] = clamp(wa(x3, t0) >> 17);
out[4] = clamp(ws(x3, t0) >> 17);
out[5] = clamp(ws(x2, t1) >> 17);
out[6] = clamp(ws(x1, t2) >> 17);
out[7] = clamp(ws(x0, t3) >> 17);
i += 8;
pos += stride;
}
}
#[inline]
#[allow(clippy::cast_possible_truncation)]
/// Multiply a number by 4096
fn f2f(x: f32) -> i32 {
(x * 4096.0 + 0.5) as i32
}
#[inline]
/// Multiply a number by 4096
fn fsh(x: i32) -> i32 {
x << 12
}
/// Clamp values between 0 and 255
#[inline]
#[allow(clippy::cast_possible_truncation)]
fn clamp(a: i32) -> i16 {
a.clamp(0, 255) as i16
}
/// IDCT assuming only the upper 4x4 is filled.
pub fn idct4x4(in_vector: &mut [i32; 64], out_vector: &mut [i16], stride: usize) {
let mut pos = 0;
// vertical pass
for ptr in 0..4 {
let i0 = wa(fsh(in_vector[ptr]), 512);
let i2 = in_vector[ptr + 16];
let p1 = wm(i2, 2217);
let p3 = wm(i2, 5352);
let x0 = wa(i0, p3);
let x1 = wa(i0, p1);
let x2 = ws(i0, p1);
let x3 = ws(i0, p3);
// odd part
let i4 = in_vector[ptr + 24];
let i3 = in_vector[ptr + 8];
let p5 = wm(wa(i4, i3), 4816);
let p1 = wa(p5, wm(i3, -3685));
let p2 = wa(p5, wm(i4, -10497));
let t3 = wa(p5, wm(i3, 867));
let t2 = wa(p5, wm(i4, -5945));
let t1 = wa(p2, wm(i3, -1597));
let t0 = wa(p1, wm(i4, -8034));
in_vector[ptr] = wa(x0, t3) >> 10;
in_vector[ptr + 8] = wa(x1, t2) >> 10;
in_vector[ptr + 16] = wa(x2, t1) >> 10;
in_vector[ptr + 24] = wa(x3, t0) >> 10;
in_vector[ptr + 32] = ws(x3, t0) >> 10;
in_vector[ptr + 40] = ws(x2, t1) >> 10;
in_vector[ptr + 48] = ws(x1, t2) >> 10;
in_vector[ptr + 56] = ws(x0, t3) >> 10;
}
// horizontal pass
for i in (0..8).map(|i| 8 * i) {
let i2 = in_vector[i + 2];
let i0 = in_vector[i];
let t0 = wa(fsh(i0), SCALE_BITS);
let t2 = wm(i2, 2217);
let t3 = wm(i2, 5352);
let x0 = wa(t0, t3);
let x3 = ws(t0, t3);
let x1 = wa(t0, t2);
let x2 = ws(t0, t2);
// odd part
let i3 = in_vector[i + 3];
let i1 = in_vector[i + 1];
let p5 = wm(wa(i3, i1), f2f(1.175875602));
let p1 = wa(p5, wm(i1, -3685));
let p2 = wa(p5, wm(i3, -10497));
let t3 = wa(p5, wm(i1, 867));
let t2 = wa(p5, wm(i3, -5945));
let t1 = wa(p2, wm(i1, -1597));
let t0 = wa(p1, wm(i3, -8034));
let out: &mut [i16; 8] = out_vector
.get_mut(pos..pos + 8)
.unwrap()
.try_into()
.unwrap();
out.copy_from_slice(&[
clamp(wa(x0, t3) >> 17),
clamp(wa(x1, t2) >> 17),
clamp(wa(x2, t1) >> 17),
clamp(wa(x3, t0) >> 17),
clamp(ws(x3, t0) >> 17),
clamp(ws(x2, t1) >> 17),
clamp(ws(x1, t2) >> 17),
clamp(ws(x0, t3) >> 17),
]);
pos += stride;
}
in_vector[32..36].fill(0);
in_vector[40..44].fill(0);
in_vector[48..52].fill(0);
in_vector[56..60].fill(0);
}
+194
View File
@@ -0,0 +1,194 @@
/*
* Copyright (c) 2023.
*
* This software is free software;
*
* You can redistribute it or modify it under terms of the MIT, Apache License or Zlib license
*/
//!This crate provides a library for decoding valid
//! ITU-T Rec. T.851 (09/2005) ITU-T T.81 (JPEG-1) or JPEG images.
//!
//!
//!
//! # Features
//! - SSE and AVX accelerated functions to speed up certain decoding operations
//! - FAST and accurate 32 bit IDCT algorithm
//! - Fast color convert functions
//! - RGBA and RGBX (4-Channel) color conversion functions
//! - YCbCr to Luma(Grayscale) conversion.
//!
//! # Usage
//! Add zune-jpeg to the dependencies in the project Cargo.toml
//!
//! ```toml
//! [dependencies]
//! zune_jpeg = "0.5"
//! ```
//! # Examples
//!
//! ## Decode a JPEG file with default arguments.
//!```no_run
//! use std::fs::read;
//! use std::io::BufReader;
//! use zune_jpeg::JpegDecoder;
//! let file_contents = BufReader::new(std::fs::File::open("a_jpeg.file").unwrap());
//! let mut decoder = JpegDecoder::new(file_contents);
//! let mut pixels = decoder.decode().unwrap();
//! ```
//!
//! ## Migrating from version 0.4--
//!
//! ### Motivation
//! zune v 0.5 reworks mainly the internal architecture of how we perform I/O
//! ,before the decoder accepted byte slices that represent the whole data as contiguous
//! but that was not ideal for all use cases, increasing memory e.g on massive files that had
//! to be read to memory.
//!
//! With v 0.5 a new I/O system is introduced, which generally introduces mechanisms to process
//! `std::io::Read + std::io::Seek` type of data feeds, (but which works in no-std), which means...
//!
//! ### What changes
//!
//! I/O code that looked like this
//!
//!```ignore
//! use zune_core::colorspace::ColorSpace;
//! use zune_jpeg::JpegDecoder;
//! // Read file into memory
//! let image = std::fs::read("image.jpg").unwrap();
//! // Make a decoder from the slice
//! let mut decoder = JpegDecoder::new(&image);
//! // decode
//! decoder.decode().unwrap();
//! ```
//!
//! Now can be rewritten in two ways.
//!
//! 1. File I/O (Using bufreader)
//!
//!```no_run
//! use std::io::BufReader;
//! use zune_core::colorspace::ColorSpace;
//! use zune_jpeg::JpegDecoder;
//!
//! let image = BufReader::new(std::fs::File::open("image.jpg").unwrap());
//! let mut decoder = JpegDecoder::new(image);
//! // decode
//! decoder.decode().unwrap();
//! ```
//!
//! 2. Reading to memory (but wrapping it in a Cursor like object)
//!```no_run
//! use zune_core::bytestream::ZCursor;
//! use zune_jpeg::JpegDecoder;
//!
//! let image_data =std::fs::read("image.jpg").unwrap();
//! // Alternatively, you can use std::io::Cursor,
//! // but it is better speed wise to use ZCursor, and it also works in
//! // no-std environments
//! let mut cursor = ZCursor::new(image_data);
//! // use the wrapped item
//! let mut decoder = JpegDecoder::new(cursor);
//! // decode
//! decoder.decode().unwrap();
//! ```
//!
//! 3. Anything that implements [ZByteReaderTrait](zune_core::bytestream::traits::ZByteReaderTrait)
//!
//! ## Decode a JPEG file to RGBA format
//!
//! - Other (limited) supported formats are and BGR, BGRA
//!
//!```no_run
//! use zune_core::bytestream::ZCursor;
//! use zune_core::colorspace::ColorSpace;
//! use zune_core::options::DecoderOptions;
//! use zune_jpeg::JpegDecoder;
//!
//! let mut options = DecoderOptions::default().jpeg_set_out_colorspace(ColorSpace::RGBA);
//!
//! let mut decoder = JpegDecoder::new_with_options(ZCursor::new(&[]),options);
//! let pixels = decoder.decode().unwrap();
//! ```
//!
//! ## Decode an image and get its width and height.
//!```no_run
//! use zune_core::bytestream::ZCursor;
//! use zune_jpeg::JpegDecoder;
//!
//! let mut decoder = JpegDecoder::new(ZCursor::new(&[]));
//! decoder.decode_headers().unwrap();
//! let image_info = decoder.info().unwrap();
//! println!("{},{}",image_info.width,image_info.height)
//! ```
//! # Crate features.
//! This crate tries to be as minimal as possible while being extensible
//! enough to handle the complexities arising from parsing different types
//! of jpeg images.
//!
//! Safety is a top concern that is why we provide both static ways to disable unsafe code,
//! disabling x86 feature, and dynamic ,by using [`DecoderOptions::set_use_unsafe(false)`],
//! both of these disable platform specific optimizations, which reduce the speed of decompression.
//!
//! Please do note that careful consideration has been taken to ensure that the unsafe paths
//! are only unsafe because they depend on platform specific intrinsics, hence no need to disable them
//!
//! The crate tries to decode as many images as possible, as a best effort, even those violating the standard
//! , this means a lot of images may get silent warnings and wrong output, but if you are sure you will be handling
//! images that follow the spec, set `ZuneJpegOptions::set_strict` to true.
//!
//![`DecoderOptions::set_use_unsafe(false)`]: https://docs.rs/zune-core/latest/zune_core/options/struct.DecoderOptions.html#method.set_use_unsafe
#![warn(
clippy::correctness,
clippy::perf,
clippy::pedantic,
clippy::inline_always,
clippy::missing_errors_doc,
clippy::panic
)]
#![allow(
clippy::needless_return,
clippy::similar_names,
clippy::inline_always,
clippy::similar_names,
clippy::doc_markdown,
clippy::module_name_repetitions,
clippy::missing_panics_doc,
clippy::missing_errors_doc
)]
// no_std compatibility
#![deny(clippy::std_instead_of_alloc, clippy::alloc_instead_of_core)]
#![cfg_attr(not(any(feature = "x86", feature = "neon")), forbid(unsafe_code))]
#![cfg_attr(not(feature = "std"), no_std)]
#![cfg_attr(feature = "portable_simd", feature(portable_simd))]
#![macro_use]
extern crate alloc;
extern crate core;
pub use zune_core;
pub use crate::components::SampleRatios;
pub use crate::decoder::{ImageInfo, JpegDecoder};
pub use crate::marker::Marker;
mod bitstream;
mod color_convert;
mod components;
mod decoder;
pub mod errors;
mod headers;
mod huffman;
#[cfg(not(fuzzing))]
mod idct;
#[cfg(fuzzing)]
pub mod idct;
mod marker;
mod mcu;
mod mcu_prog;
mod misc;
mod unsafe_utils;
mod unsafe_utils_avx2;
mod unsafe_utils_neon;
mod upsampler;
mod worker;
+91
View File
@@ -0,0 +1,91 @@
/*
* Copyright (c) 2023.
*
* This software is free software;
*
* You can redistribute it or modify it under terms of the MIT, Apache License or Zlib license
*/
#![allow(clippy::upper_case_acronyms)]
/// JPEG Markers
///
/// **NOTE** This doesn't cover all markers, just the ones zune-jpeg supports.
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum Marker {
/// Start Of Frame markers
///
/// - SOF(0): Baseline DCT (Huffman coding)
/// - SOF(1): Extended sequential DCT (Huffman coding)
/// - SOF(2): Progressive DCT (Huffman coding)
/// - SOF(3): Lossless (sequential) (Huffman coding)
/// - SOF(5): Differential sequential DCT (Huffman coding)
/// - SOF(6): Differential progressive DCT (Huffman coding)
/// - SOF(7): Differential lossless (sequential) (Huffman coding)
/// - SOF(9): Extended sequential DCT (arithmetic coding)
/// - SOF(10): Progressive DCT (arithmetic coding)
/// - SOF(11): Lossless (sequential) (arithmetic coding)
/// - SOF(13): Differential sequential DCT (arithmetic coding)
/// - SOF(14): Differential progressive DCT (arithmetic coding)
/// - SOF(15): Differential lossless (sequential) (arithmetic coding)
SOF(u8),
/// Define Huffman table(s)
DHT,
/// Define arithmetic coding conditioning(s)
DAC,
/// Restart with modulo 8 count `m`
RST(u8),
/// Start of image
SOI,
/// End of image
EOI,
/// Start of scan
SOS,
/// Define quantization table(s)
DQT,
/// Define number of lines
DNL,
/// Define restart interval
DRI,
/// Reserved for application segments
APP(u8),
/// Comment
COM,
/// Unknown markers
UNKNOWN(u8)
}
impl Marker {
pub fn from_u8(n: u8) -> Option<Marker> {
use self::Marker::{APP, COM, DAC, DHT, DNL, DQT, DRI, EOI, RST, SOF, SOI, SOS, UNKNOWN};
match n {
0xFE => Some(COM),
0xC0 => Some(SOF(0)),
0xC1 => Some(SOF(1)),
0xC2 => Some(SOF(2)),
0xC4 => Some(DHT),
0xCC => Some(DAC),
0xD0 => Some(RST(0)),
0xD1 => Some(RST(1)),
0xD2 => Some(RST(2)),
0xD3 => Some(RST(3)),
0xD4 => Some(RST(4)),
0xD5 => Some(RST(5)),
0xD6 => Some(RST(6)),
0xD7 => Some(RST(7)),
0xD8 => Some(SOI),
0xD9 => Some(EOI),
0xDA => Some(SOS),
0xDB => Some(DQT),
0xDC => Some(DNL),
0xDD => Some(DRI),
0xE0 => Some(APP(0)),
0xE1 => Some(APP(1)),
0xE2 => Some(APP(2)),
0xED => Some(APP(13)),
0xEE => Some(APP(14)),
_ => Some(UNKNOWN(n))
}
}
}
+936
View File
@@ -0,0 +1,936 @@
/*
* Copyright (c) 2023.
*
* This software is free software;
*
* You can redistribute it or modify it under terms of the MIT, Apache License or Zlib license
*/
use alloc::vec::Vec;
use alloc::{format, vec};
use core::cmp::min;
use zune_core::bytestream::ZByteReaderTrait;
use zune_core::colorspace::ColorSpace;
use zune_core::colorspace::ColorSpace::Luma;
use zune_core::log::{error, trace, warn};
use crate::bitstream::BitStream;
use crate::components::SampleRatios;
use crate::decoder::MAX_COMPONENTS;
use crate::errors::DecodeErrors;
use crate::marker::Marker;
use crate::mcu_prog::get_marker;
use crate::misc::{calculate_padded_width, setup_component_params};
use crate::worker::{color_convert, upsample};
use crate::JpegDecoder;
/// The size of a DC block for a MCU.
pub const DCT_BLOCK: usize = 64;
impl<T: ZByteReaderTrait> JpegDecoder<T> {
/// Check for existence of DC and AC Huffman Tables
pub(crate) fn check_tables(&self) -> Result<(), DecodeErrors> {
// check that dc and AC tables exist outside the hot path
for component in &self.components {
let _ = &self
.dc_huffman_tables
.get(component.dc_huff_table)
.as_ref()
.ok_or_else(|| {
DecodeErrors::HuffmanDecode(format!(
"No Huffman DC table for component {:?} ",
component.component_id
))
})?
.as_ref()
.ok_or_else(|| {
DecodeErrors::HuffmanDecode(format!(
"No DC table for component {:?}",
component.component_id
))
})?;
let _ = &self
.ac_huffman_tables
.get(component.ac_huff_table)
.as_ref()
.ok_or_else(|| {
DecodeErrors::HuffmanDecode(format!(
"No Huffman AC table for component {:?} ",
component.component_id
))
})?
.as_ref()
.ok_or_else(|| {
DecodeErrors::HuffmanDecode(format!(
"No AC table for component {:?}",
component.component_id
))
})?;
}
Ok(())
}
/// Decode MCUs and carry out post processing.
///
/// This is the main decoder loop for the library, the hot path.
///
/// Because of this, we pull in some very crazy optimization tricks hence readability is a pinch
/// here.
#[allow(
clippy::similar_names,
clippy::too_many_lines,
clippy::cast_possible_truncation
)]
#[inline(never)]
pub(crate) fn decode_mcu_ycbcr_baseline(
&mut self, pixels: &mut [u8]
) -> Result<(), DecodeErrors> {
setup_component_params(self)?;
// check dc and AC tables
self.check_tables()?;
let (mut mcu_width, mut mcu_height);
if self.is_interleaved {
// set upsampling functions
self.set_upsampling()?;
mcu_width = self.mcu_x;
mcu_height = self.mcu_y;
} else {
// For non-interleaved images( (1*1) subsampling)
// number of MCU's are the widths (+7 to account for paddings) divided bu 8.
mcu_width = ((self.info.width + 7) / 8) as usize;
mcu_height = ((self.info.height + 7) / 8) as usize;
}
if self.is_interleaved
&& self.input_colorspace.num_components() > 1
&& self.options.jpeg_get_out_colorspace().num_components() == 1
&& (self.info.sample_ratio == SampleRatios::V
|| self.info.sample_ratio == SampleRatios::HV)
{
// For a specific set of images, e.g interleaved,
// when converting from YcbCr to grayscale, we need to
// take into account mcu height since the MCU decoding needs to take
// it into account for padding purposes and the post processor
// parses two rows per mcu width.
//
// set coeff to be 2 to ensure that we increment two rows
// for every mcu processed also
mcu_height *= self.v_max;
mcu_height /= self.h_max;
self.coeff = 2;
}
if self.input_colorspace == ColorSpace::Luma && self.is_interleaved {
warn!("Grayscale image with down-sampled component, resetting component details");
self.reset_params();
mcu_width = ((self.info.width + 7) / 8) as usize;
mcu_height = ((self.info.height + 7) / 8) as usize;
}
let width = usize::from(self.info.width);
let padded_width = calculate_padded_width(width, self.info.sample_ratio);
let mut stream = BitStream::new();
let mut tmp = [0_i32; DCT_BLOCK];
let comp_len = self.components.len();
for (pos, comp) in self.components.iter_mut().enumerate() {
// Allocate only needed components.
//
// For special colorspaces i.e YCCK and CMYK, just allocate all of the needed
// components.
if min(
self.options.jpeg_get_out_colorspace().num_components() - 1,
pos
) == pos
|| comp_len == 4
// Special colorspace
{
// allocate enough space to hold a whole MCU width
// this means we should take into account sampling ratios
// `*8` is because each MCU spans 8 widths.
let len = comp.width_stride * comp.vertical_sample * 8;
comp.needed = true;
comp.raw_coeff = vec![0; len];
} else {
comp.needed = false;
}
}
// If all components are contained in the first scan of MCUs, then we can process into
// (upsampled) pixels immediately after each MCU, for convenience we use each row of MCUS.
// Otherwise, we must first wait until following SOS provide the remaining components.
let all_components_in_first_scan = usize::from(self.num_scans) == self.components.len();
let mut progressive_mcus: [Vec<i16>; 4] = core::array::from_fn(|_| vec![]);
if !all_components_in_first_scan {
for (component, mcu) in self.components.iter().zip(&mut progressive_mcus) {
let len = mcu_width
* component.vertical_sample
* component.horizontal_sample
* mcu_height
* 64;
*mcu = vec![0; len];
}
}
let mut pixels_written = 0;
let is_hv = usize::from(self.is_interleaved);
let upsampler_scratch_size = is_hv * self.components.iter().map(|x| x.width_stride).max().unwrap_or(0) * 8;
let mut upsampler_scratch_space = vec![0; upsampler_scratch_size];
'sos: loop {
trace!(
"Baseline decoding of components: {:?}",
&self.z_order[..usize::from(self.num_scans)]
);
trace!("Decoding MCU width: {mcu_width}, height: {mcu_height}");
for i in 0..mcu_height {
if stream.overread_by > 0 {
pixels.get_mut(pixels_written..).map(|v| v.fill(128));
if self.options.strict_mode() {
return Err(DecodeErrors::FormatStatic("Premature end of buffer"));
};
error!("Premature end of buffer");
break;
}
// decode a whole MCU width,
// this takes into account interleaved components.
let terminate = if all_components_in_first_scan {
self.decode_mcu_width::<false>(
mcu_width,
i,
&mut tmp,
&mut stream,
&mut progressive_mcus
)?
} else {
/* NB: (cae). This code was added due to the issue at https://github.com/etemesi254/zune-image/issues/277
*
* There is a particular set of images that interleave the start of scan (SOS) with the MCU,
* E.g if it's a three component image, we have SOS->MCU ->SOS->MCU ->SOS->MCU
* which presents a problem on decoding, we need to buffer the whole image before continuing since
* we won't have a row containing all the component data which will be needed e.g for color conversion.
*
* The mechanisms is that we decode the whole image upfront, which goes against the normal
* routine of decoding MCU width , so this requires more memory upfront than initial routines
* but it is a single image out of the many corpuses that exist, so its fine.
* (image in test-images/jpeg/sos_news.jpeg)
* Code contributed by Aurelia Molzer (https://github.com/197g)
*
*/
self.decode_mcu_width::<true>(
mcu_width,
i,
&mut tmp,
&mut stream,
&mut progressive_mcus
)?
};
// process that width up until it's impossible. This is faster than allocation the
// full components, which we skipped earlier.
if all_components_in_first_scan {
self.post_process(
pixels,
i,
mcu_height,
width,
padded_width,
&mut pixels_written,
&mut upsampler_scratch_space
)?;
}
match terminate {
McuContinuation::Ok => {}
McuContinuation::AnotherSos if all_components_in_first_scan => {
warn!("More than one SOS despite already having all components");
return Ok(());
}
McuContinuation::AnotherSos => continue 'sos,
McuContinuation::InterScanMarker(marker) => {
// Handle inter-scan markers (DHT/DQT/etc) uniformly here.
// This keeps all marker handling in the outer loop.
if self.advance_to_next_sos(marker, &mut stream)? {
continue 'sos;
} else {
// Hit EOI
break;
}
}
McuContinuation::Terminate => {
warn!("Got terminate signal, will not process further");
pixels.get_mut(pixels_written..).map(|v| v.fill(128));
return Ok(());
}
}
}
// Breaks if we get here, looping only if we have restarted, i.e. found another SOS and
// continued at `'sos'.
break;
}
if !all_components_in_first_scan {
self.finish_baseline_decoding(&progressive_mcus, mcu_width, pixels)?;
}
// it may happen that some images don't have the whole buffer
// so we can't panic in case of that
// assert_eq!(pixels_written, pixels.len());
// For UHD usecases that tie two images separating them with EOI and
// SOI markers, it may happen that we do not reach this image end of image
// So this ensures we reach it
// Ensure we read EOI
if !stream.seen_eoi {
let marker = get_marker(&mut self.stream, &mut stream);
match marker {
Ok(_m) => {
trace!("Found marker {:?}", _m);
}
Err(_) => {
// ignore error
}
}
}
trace!("Finished decoding image");
Ok(())
}
/// Process all MCUs when baseline decoding has been processing them component-after-component.
/// For simplicity this assembles the dequantized blocks in the order that the post processing
/// of an interleaved baseline decoding would use.
#[allow(clippy::too_many_lines)]
#[allow(clippy::cast_sign_loss)]
pub(crate) fn finish_baseline_decoding(
&mut self, block: &[Vec<i16>; MAX_COMPONENTS], _mcu_width: usize, pixels: &mut [u8]
) -> Result<(), DecodeErrors> {
let mcu_height = self.mcu_y;
// Size of our output image(width*height)
let is_hv = usize::from(self.is_interleaved);
let upsampler_scratch_size = is_hv * self.components[0].width_stride;
let width = usize::from(self.info.width);
let padded_width = calculate_padded_width(width, self.info.sample_ratio);
let mut upsampler_scratch_space = vec![0; upsampler_scratch_size];
for (pos, comp) in self.components.iter_mut().enumerate() {
// Mark only needed components for computing output colors.
if min(
self.options.jpeg_get_out_colorspace().num_components() - 1,
pos
) == pos
|| self.input_colorspace == ColorSpace::YCCK
|| self.input_colorspace == ColorSpace::CMYK
{
comp.needed = true;
} else {
comp.needed = false;
}
}
let mut pixels_written = 0;
// dequantize and idct have been performed, only color convert.
for i in 0..mcu_height {
// All the data is already in the right order, we just need to be able to pass it to
// the post_process & upsample method. That expects all the data to be stored as one
// row of MCUs in each component's `raw_coeff`.
'component: for (position, component) in &mut self.components.iter_mut().enumerate() {
if !component.needed {
continue 'component;
}
// step is the number of pixels this iteration wil be handling
// Given by the number of mcu's height and the length of the component block
// Since the component block contains the whole channel as raw pixels
// we this evenly divides the pixels into MCU blocks
//
// For interleaved images, this gives us the exact pixels comprising a whole MCU
// block
let step = block[position].len() / mcu_height;
// where we will be reading our pixels from.
let slice = &block[position][i * step..][..step];
let temp_channel = &mut component.raw_coeff;
temp_channel[..step].copy_from_slice(slice);
}
// process that whole stripe of MCUs
self.post_process(
pixels,
i,
mcu_height,
width,
padded_width,
&mut pixels_written,
&mut upsampler_scratch_space
)?;
}
return Ok(());
}
fn decode_mcu_width<const PROGRESSIVE: bool>(
&mut self, mcu_width: usize, mcu_height: usize, tmp: &mut [i32; 64],
stream: &mut BitStream, progressive: &mut [Vec<i16>; 4]
) -> Result<McuContinuation, DecodeErrors> {
let is_one_by_one = !self.scan_subsampled;
// The definition of MCU depends on the sampling factor of involved scans. When components
// have different factors then each Minimal-Coding-Unit is the least common multiple such
// that we have an integer number of blocks from each component. But the decoding of these
// components differs from it otherwise, we need an inner loop with a dynamic amount of
// coefficients per component, whereas otherwise we have exactly one block of coefficients
// encoded for each component in the bitstream order.
//
// We statically specialize on this to improve code generation of the common case a little
// bit. We could also special case common sub-sampling cases but be mindful of code bloat.
if is_one_by_one {
self.inner_decode_mcu_width::<PROGRESSIVE, false>(
mcu_width,
mcu_height,
tmp,
stream,
progressive
)
} else {
self.inner_decode_mcu_width::<PROGRESSIVE, true>(
mcu_width,
mcu_height,
tmp,
stream,
progressive
)
}
}
// Inline-never ensures we do get this function optimize on its own, into two different
// versions, without the optimizer tripping up over the complexity that comes with the
// constant folding. And constant folding is quite important for performance here as
// when `not SAMPLED` then the inner loop has exactly one iteration per component in
// the scan. The difference was ~1% or a bit more.
fn inner_decode_mcu_width<const PROGRESSIVE: bool, const SAMPLED: bool>(
&mut self, mcu_width: usize, mcu_height: usize, tmp: &mut [i32; 64],
stream: &mut BitStream, progressive: &mut [Vec<i16>; 4]
) -> Result<McuContinuation, DecodeErrors> {
let z_order = self.z_order;
let z_scans = &z_order[..usize::from(self.num_scans)];
// How much of the head of `tmp` was written by the last MCU decoding? We only check for
// two different cases and not all possible outcomes as this is only used to optimize the
// bytes written in `fill`. Since the clobber happens in UNZIGZAG order we'd be straddling
// most cache lines anyways even if we did a partial write with the exact length of the
// coefficient data which was written into `tmp`.
let mut clobber_more_than_4x4 = true;
// For non-interleaved scans (PROGRESSIVE=true), each scan contains a single component
// and we iterate over that component's actual data unit count, not the interleaved MCU
// width multiplied by sampling factor.
let scan_du_width = if PROGRESSIVE {
let k = z_scans[0];
let comp = &self.components[k];
// Calculate actual data units for this component: ceil(width / (8 * subsampling_ratio))
(self.info.width as usize * comp.horizontal_sample + self.h_max * 8 - 1)
/ (self.h_max * 8)
} else {
mcu_width
};
for j in 0..scan_du_width {
// iterate over components
for &k in z_scans {
// we made this loop body massive due to several different paths that depend on
// static conditions. Note we (potentially) call into other functions so the
// compiler will not unroll anything here anyways. The gains from separating
// differently optimized loop bodies are much greater than a single additional jump
// here.
let component = &mut self.components[k];
let dc_table = self.dc_huffman_tables[component.dc_huff_table % MAX_COMPONENTS]
.as_ref()
.ok_or(DecodeErrors::FormatStatic("DC table not found"))?;
let ac_table = self.ac_huffman_tables[component.ac_huff_table % MAX_COMPONENTS]
.as_ref()
.ok_or(DecodeErrors::FormatStatic("AC table not found"))?;
let qt_table = &component.quantization_table;
let channel = if PROGRESSIVE {
let offset =
mcu_height * component.width_stride * 8 * component.vertical_sample;
&mut progressive[k][offset..]
} else {
&mut component.raw_coeff
};
let component_samples_needed = component.needed;
// If image is interleaved iterate over scan components,
// otherwise if it-s non-interleaved, these routines iterate in
// trivial scanline order(Y,Cb,Cr)
//
// Turn the bounds into a compile time constant for a common special case. This
// allows the compiler to unroll the loop and then do a bunch of interleaving.
//
// For PROGRESSIVE (non-interleaved), we iterate data units directly so
// h_samp/v_samp loops run exactly once.
let v_step =
if SAMPLED && !PROGRESSIVE { 0..component.vertical_sample } else { 0..1 };
for v_samp in v_step {
let h_step =
if SAMPLED && !PROGRESSIVE { 0..component.horizontal_sample } else { 0..1 };
for h_samp in h_step {
let result = if component_samples_needed {
// Fill the array with zeroes, decode_mcu_block expects
// a zero based array. Clobber is in zig-zag order though.
// Writing consecutive entries is basically free in terms
// of memory throughput so we opt for a larger power of
// two which lets the compiler turn this into a repeated
// write of a zeroed vector register, which does not have
// any branches, instead of a more difficult pattern where
// we attempt to overwrite exactly one coefficient.
let clobber_len = if !clobber_more_than_4x4 { 32 } else { 64 };
tmp[..clobber_len].fill(0);
stream.decode_mcu_block(
&mut self.stream,
dc_table,
ac_table,
qt_table,
tmp,
&mut component.dc_pred
)
} else {
// We do not touch tmp so there is no need to reset it.
stream.discard_mcu_block(&mut self.stream, dc_table, ac_table)
};
// If an error occurs we can either propagate it
// as an error or print it and call terminate.
//
// This allows even corrupt images to render something,
// even if its bad, matching browsers.
//
// See example in https://github.com/etemesi254/zune-image/issues/293
let len = if let Ok(len) = result {
len
} else {
// result.is_err()
return if self.options.strict_mode() {
Err(result.err().unwrap())
} else {
error!("{}", result.err().unwrap());
Ok(McuContinuation::Terminate)
};
};
if component_samples_needed {
// tmp was only written partially, note that len is in ZigZag order.
clobber_more_than_4x4 = len > 10;
let idct_position = if PROGRESSIVE {
// For non-interleaved, j indexes data units directly
j * 8
} else {
// derived from stb and rewritten for my tastes
let c2 = v_samp * 8;
let c3 = ((j * component.horizontal_sample) + h_samp) * 8;
component.width_stride * c2 + c3
};
let idct_pos = channel.get_mut(idct_position..).unwrap();
if len <= 1 {
(self.idct_1x1_func)(tmp, idct_pos, component.width_stride);
} else if len <= 10 {
(self.idct_4x4_func)(tmp, idct_pos, component.width_stride);
} else {
// call idct.
(self.idct_func)(tmp, idct_pos, component.width_stride);
}
}
}
}
}
self.todo = self.todo.wrapping_sub(1);
if self.todo == 0 {
self.handle_rst_main(stream)?;
continue;
}
if stream.marker.is_some() && stream.bits_left == 0 {
break;
}
}
self.check_stream_marker_after_mcu_width(stream)
}
fn check_stream_marker_after_mcu_width(
&mut self, stream: &mut BitStream
) -> Result<McuContinuation, DecodeErrors> {
// After all interleaved components, that's an MCU
// handle stream markers
//
// In some corrupt images, it may occur that header markers occur in the stream.
// The spec EXPLICITLY FORBIDS this, specifically, in
// routine F.2.2.5 it says
// `The only valid marker which may occur within the Huffman coded data is the RSTm marker.`
//
// But libjpeg-turbo allows it because of some weird reason. so I'll also
// allow it because of some weird reason.
if let Some(m) = stream.marker {
if m == Marker::EOI {
// acknowledge and ignore EOI marker.
stream.marker.take();
trace!("Found EOI marker");
// Google Introduced the Ultra-HD image format which is basically
// stitching two images into one container.
// They basically separate two images via a EOI and SOI marker
// so let's just ensure if we ever see EOI, we never read past that
// ever.
// https://github.com/google/libultrahdr
stream.seen_eoi = true;
} else if let Marker::RST(_) = m {
//debug_assert_eq!(self.todo, 0);
if self.todo == 0 {
self.handle_rst(stream)?;
}
} else if let Marker::SOS = m {
self.parse_marker_inner(m)?;
stream.marker.take();
stream.reset();
trace!("Found SOS marker");
return Ok(McuContinuation::AnotherSos);
} else if matches!(m, Marker::DHT | Marker::DQT | Marker::DRI | Marker::COM)
|| matches!(m, Marker::APP(_))
{
// For non-interleaved images, setup markers can appear between scans.
// Signal the caller to handle this marker and find the next SOS.
// This keeps all marker parsing in the caller's loop.
stream.marker.take();
trace!("Found inter-scan marker {:?}", m);
return Ok(McuContinuation::InterScanMarker(m));
} else {
if self.options.strict_mode() {
return Err(DecodeErrors::Format(format!(
"Marker {m:?} found where not expected"
)));
}
error!(
"Marker `{:?}` Found within Huffman Stream, possibly corrupt jpeg",
m
);
self.parse_marker_inner(m)?;
stream.marker.take();
stream.reset();
return Ok(McuContinuation::Terminate);
}
}
Ok(McuContinuation::Ok)
}
/// Scan for the next SOS marker, parsing setup markers along the way.
///
/// This is the unified marker scanning function used after encountering an
/// inter-scan marker. It handles DHT, DQT, DRI, COM, and APP markers that
/// can appear between scans in non-interleaved images.
///
/// # Arguments
/// * `first_marker` - The first marker that was already detected (not yet parsed)
/// * `stream` - The bitstream state
///
/// # Returns
/// * `Ok(true)` - Found SOS, ready to continue decoding
/// * `Ok(false)` - Found EOI, decoding complete
/// * `Err(_)` - Error (too many markers, unexpected marker in strict mode, etc.)
fn advance_to_next_sos(
&mut self,
first_marker: Marker,
stream: &mut BitStream
) -> Result<bool, DecodeErrors> {
// Limit iterations to prevent DoS from malicious files.
const MAX_INTER_SCAN_MARKERS: usize = 64;
// Parse the first marker that triggered this call
self.parse_marker_inner(first_marker)?;
stream.reset();
for _ in 0..MAX_INTER_SCAN_MARKERS {
let marker = get_marker(&mut self.stream, stream)?;
match marker {
Marker::SOS => {
self.parse_marker_inner(Marker::SOS)?;
stream.reset();
trace!("Found SOS marker, continuing decode");
return Ok(true);
}
Marker::EOI => {
stream.seen_eoi = true;
trace!("Found EOI marker");
return Ok(false);
}
Marker::DHT | Marker::DQT | Marker::DRI | Marker::COM => {
trace!("Parsing inter-scan marker {:?}", marker);
self.parse_marker_inner(marker)?;
}
Marker::APP(_) => {
trace!("Parsing inter-scan APP marker {:?}", marker);
self.parse_marker_inner(marker)?;
}
other => {
if self.options.strict_mode() {
return Err(DecodeErrors::Format(format!(
"Unexpected marker {:?} while scanning for SOS between scans",
other
)));
}
// Non-strict: skip unknown marker
warn!("Skipping unexpected marker {:?} between scans", other);
let length = self.stream.get_u16_be_err()?;
if length >= 2 {
self.stream.skip((length - 2) as usize)?;
}
}
}
}
Err(DecodeErrors::FormatStatic(
"Too many markers between scans (exceeded limit of 64)"
))
}
// handle RST markers.
// No-op if not using restarts
// this routine is shared with mcu_prog
#[cold]
pub(crate) fn handle_rst(&mut self, stream: &mut BitStream) -> Result<(), DecodeErrors> {
self.todo = self.restart_interval;
if let Some(marker) = stream.marker {
// Found a marker
// Read stream and see what marker is stored there
match marker {
Marker::RST(_) => {
// reset stream
stream.reset();
// Initialize dc predictions to zero for all components
self.components.iter_mut().for_each(|x| x.dc_pred = 0);
// Start iterating again. from position.
}
Marker::EOI => {
// silent pass
}
_ => {
return Err(DecodeErrors::MCUError(format!(
"Marker {marker:?} found in bitstream, possibly corrupt jpeg"
)));
}
}
}
Ok(())
}
#[allow(clippy::too_many_lines, clippy::too_many_arguments)]
pub(crate) fn post_process(
&mut self, pixels: &mut [u8], i: usize, mcu_height: usize, width: usize,
padded_width: usize, pixels_written: &mut usize, upsampler_scratch_space: &mut [i16]
) -> Result<(), DecodeErrors> {
let out_colorspace_components = self.options.jpeg_get_out_colorspace().num_components();
let mut px = *pixels_written;
// indicates whether image is vertically up-sampled
let is_vertically_sampled = self
.components
.iter()
.any(|c| c.sample_ratio == SampleRatios::HV || c.sample_ratio == SampleRatios::V);
let mut comp_len = self.components.len();
// If we are moving from YCbCr -> Luma, we do not allocate storage for other components, so we
// will panic when we are trying to read samples, so for that case,
// hardcode it so that we don't panic when doing
// *samp = &samples[j][pos * padded_width..(pos + 1) * padded_width]
if out_colorspace_components < comp_len && self.options.jpeg_get_out_colorspace() == Luma {
comp_len = out_colorspace_components;
}
let mut color_conv_function =
|num_iters: usize, samples: [&[i16]; 4]| -> Result<(), DecodeErrors> {
for (pos, output) in pixels[px..]
.chunks_exact_mut(width * out_colorspace_components)
.take(num_iters)
.enumerate()
{
let mut raw_samples: [&[i16]; 4] = [&[], &[], &[], &[]];
// iterate over each line, since color-convert needs only
// one line
for (j, samp) in raw_samples.iter_mut().enumerate().take(comp_len) {
let temp = &samples[j].get(pos * padded_width..(pos + 1) * padded_width);
if temp.is_none() {
return Err(DecodeErrors::FormatStatic("Missing samples"));
}
*samp = temp.unwrap();
}
color_convert(
&raw_samples,
self.color_convert_16,
self.input_colorspace,
self.options.jpeg_get_out_colorspace(),
output,
width,
padded_width
)?;
px += width * out_colorspace_components;
}
Ok(())
};
let comps = &mut self.components[..];
if self.is_interleaved && self.options.jpeg_get_out_colorspace() != ColorSpace::Luma {
for comp in comps.iter_mut() {
upsample(
comp,
mcu_height,
i,
upsampler_scratch_space,
is_vertically_sampled
)?;
}
if is_vertically_sampled {
if i > 0 {
// write the last line, it wasn't up-sampled as we didn't have row_down
// yet
let mut samples: [&[i16]; 4] = [&[], &[], &[], &[]];
for (samp, component) in samples.iter_mut().zip(comps.iter()) {
*samp = &component.first_row_upsample_dest;
}
// ensure length matches for all samples
let _first_len = samples[0].len();
// This was a good check, but can be caused to panic, esp on invalid/corrupt images.
// See one in issue https://github.com/etemesi254/zune-image/issues/262, so for now
// we just ignore and generate invalid images at the end.
//
//
// for samp in samples.iter().take(comp_len) {
// assert_eq!(first_len, samp.len());
// }
let num_iters = self.coeff * self.v_max;
color_conv_function(num_iters, samples)?;
}
// After up-sampling the last row, save any row that can be used for
// a later up-sampling,
//
// E.g the Y sample is not sampled but we haven't finished upsampling the last row of
// the previous mcu, since we don't have the down row, so save it
for component in comps.iter_mut() {
if component.sample_ratio != SampleRatios::H {
// We don't care about H sampling factors, since it's copied in the workers function
// copy last row to be used for the next color conversion
let size = component.vertical_sample
* component.width_stride
* component.sample_ratio.sample();
let last_bytes =
component.raw_coeff.rchunks_exact_mut(size).next().unwrap();
component
.first_row_upsample_dest
.copy_from_slice(last_bytes);
}
}
}
let mut samples: [&[i16]; 4] = [&[], &[], &[], &[]];
for (samp, component) in samples.iter_mut().zip(comps.iter()) {
*samp = if component.sample_ratio == SampleRatios::None {
&component.raw_coeff
} else {
&component.upsample_dest
};
}
// we either do 7 or 8 MCU's depending on the state, this only applies to
// vertically sampled images
//
// for rows up until the last MCU, we do not upsample the last stride of the MCU
// which means that the number of iterations should take that into account is one less the
// up-sampled size
//
// For the last MCU, we upsample the last stride, meaning that if we hit the last MCU, we
// should sample full raw coeffs
let is_last_considered = is_vertically_sampled && (i != mcu_height.saturating_sub(1));
let num_iters = (8 - usize::from(is_last_considered)) * self.coeff * self.v_max;
color_conv_function(num_iters, samples)?;
} else {
let mut channels_ref: [&[i16]; MAX_COMPONENTS] = [&[]; MAX_COMPONENTS];
self.components
.iter()
.enumerate()
.for_each(|(pos, x)| channels_ref[pos] = &x.raw_coeff);
if let SampleRatios::Generic(_, v) = self.info.sample_ratio {
color_conv_function(8 * v * self.coeff, channels_ref)?;
} else {
color_conv_function(8 * self.coeff, channels_ref)?;
}
}
*pixels_written = px;
Ok(())
}
}
enum McuContinuation {
Ok,
AnotherSos,
/// Found an inter-scan marker (DHT/DQT/DRI/COM/APP) that needs handling.
/// The caller should parse it and scan for the next SOS.
InterScanMarker(Marker),
Terminate
}
+688
View File
@@ -0,0 +1,688 @@
/*
* Copyright (c) 2023.
*
* This software is free software;
*
* You can redistribute it or modify it under terms of the MIT, Apache License or Zlib license
*/
//!Routines for progressive decoding
/*
This file is needlessly complicated,
It is that way to ensure we don't burn memory anyhow
Memory is a scarce resource in some environments, I would like this to be viable
in such environments
Half of the complexity comes from the jpeg spec, because progressive decoding,
is one hell of a ride.
*/
use alloc::string::ToString;
use alloc::vec::Vec;
use alloc::{format, vec};
use core::cmp::min;
use zune_core::bytestream::{ZByteReaderTrait, ZReader};
use zune_core::colorspace::ColorSpace;
use zune_core::log::{debug, error, warn};
use crate::bitstream::BitStream;
use crate::components::SampleRatios;
use crate::decoder::{JpegDecoder, MAX_COMPONENTS};
use crate::errors::DecodeErrors;
use crate::headers::parse_sos;
use crate::marker::Marker;
use crate::mcu::DCT_BLOCK;
use crate::misc::{calculate_padded_width, setup_component_params};
impl<T: ZByteReaderTrait> JpegDecoder<T> {
/// Decode a progressive image
///
/// This routine decodes a progressive image, stopping if it finds any error.
#[allow(
clippy::needless_range_loop,
clippy::cast_sign_loss,
clippy::redundant_else,
clippy::too_many_lines
)]
#[inline(never)]
pub(crate) fn decode_mcu_ycbcr_progressive(
&mut self, pixels: &mut [u8]
) -> Result<(), DecodeErrors> {
setup_component_params(self)?;
let mut mcu_height;
// memory location for decoded pixels for components
let mut block: [Vec<i16>; MAX_COMPONENTS] = [vec![], vec![], vec![], vec![]];
let mut mcu_width;
let mut seen_scans = 1;
if self.input_colorspace == ColorSpace::Luma && self.is_interleaved {
warn!("Grayscale image with down-sampled component, resetting component details");
self.reset_params();
}
if self.is_interleaved {
// this helps us catch component errors.
self.set_upsampling()?;
}
if self.is_interleaved {
mcu_width = self.mcu_x;
mcu_height = self.mcu_y;
} else {
mcu_width = (self.info.width as usize + 7) / 8;
mcu_height = (self.info.height as usize + 7) / 8;
}
if self.is_interleaved
&& self.input_colorspace.num_components() > 1
&& self.options.jpeg_get_out_colorspace().num_components() == 1
&& (self.info.sample_ratio == SampleRatios::V
|| self.info.sample_ratio == SampleRatios::HV)
{
// For a specific set of images, e.g interleaved,
// when converting from YcbCr to grayscale, we need to
// take into account mcu height since the MCU decoding needs to take
// it into account for padding purposes and the post processor
// parses two rows per mcu width.
//
// set coeff to be 2 to ensure that we increment two rows
// for every mcu processed also
mcu_height *= self.v_max;
mcu_height /= self.h_max;
self.coeff = 2;
}
mcu_width *= 64;
for i in 0..self.input_colorspace.num_components() {
let comp = &self.components[i];
let len = mcu_width * comp.vertical_sample * comp.horizontal_sample * mcu_height;
block[i] = vec![0; len];
}
let mut stream = BitStream::new_progressive(self.succ_low, self.spec_start, self.spec_end);
// there are multiple scans in the stream, this should resolve the first scan
let result = self.parse_entropy_coded_data(&mut stream, &mut block);
if result.is_err() {
return if self.options.strict_mode() {
Err(result.err().unwrap())
} else {
error!("{}", result.err().unwrap());
// Go process it and return as much as we can, exiting here
return self.finish_progressive_decoding(&block, pixels);
};
}
// extract marker
let mut marker = stream
.marker
.take()
.ok_or(DecodeErrors::FormatStatic("Marker missing where expected"))?;
// if marker is EOI, we are done, otherwise continue scanning.
//
// In case we have a premature image, we print a warning or return
// an error, depending on the strictness of the decoder, so there
// is that logic to handle too
'eoi: while marker != Marker::EOI {
match marker {
Marker::SOS => {
parse_sos(self)?;
stream.update_progressive_params(
self.succ_high,
self.succ_low,
self.spec_start,
self.spec_end
);
// after every SOS, marker, parse data for that scan.
let result = self.parse_entropy_coded_data(&mut stream, &mut block);
// Do not error out too fast, allows the decoder to continue as much as possible
// even after errors
if result.is_err() {
return if self.options.strict_mode() {
Err(result.err().unwrap())
} else {
error!("{}", result.err().unwrap());
break 'eoi;
};
}
// extract marker, might either indicate end of image or we continue
// scanning(hence the continue statement to determine).
match get_marker(&mut self.stream, &mut stream) {
Ok(marker_n) => {
marker = marker_n;
seen_scans += 1;
if seen_scans > self.options.jpeg_get_max_scans() {
return Err(DecodeErrors::Format(format!(
"Too many scans, exceeded limit of {}",
self.options.jpeg_get_max_scans()
)));
}
stream.reset();
continue 'eoi;
}
Err(msg) => {
if self.options.strict_mode() {
return Err(msg);
}
error!("{:?}", msg);
break 'eoi;
}
}
}
Marker::RST(_n) => {
self.handle_rst(&mut stream)?;
}
_ => {
self.parse_marker_inner(marker)?;
}
}
match get_marker(&mut self.stream, &mut stream) {
Ok(marker_n) => {
marker = marker_n;
}
Err(e) => {
if self.options.strict_mode() {
return Err(e);
}
error!("{}", e);
// If we can't get the marker, just break away
// allows us to decode some corrupt images
// e.g https://github.com/etemesi254/zune-image/issues/294
break 'eoi;
}
}
}
self.finish_progressive_decoding(&block, pixels)
}
/// Reset progressive parameters
fn reset_prog_params(&mut self, stream: &mut BitStream) {
stream.reset();
self.components.iter_mut().for_each(|x| x.dc_pred = 0);
// Also reset JPEG restart intervals
self.todo = if self.restart_interval != 0 { self.restart_interval } else { usize::MAX };
}
#[allow(clippy::too_many_lines, clippy::cast_sign_loss)]
fn parse_entropy_coded_data(
&mut self, stream: &mut BitStream, buffer: &mut [Vec<i16>; MAX_COMPONENTS]
) -> Result<(), DecodeErrors> {
self.reset_prog_params(stream);
if usize::from(self.num_scans) > self.input_colorspace.num_components() {
return Err(DecodeErrors::Format(format!(
"Number of scans {} cannot be greater than number of components, {}",
self.num_scans,
self.input_colorspace.num_components()
)));
}
if self.num_scans == 1 {
// Safety checks
if self.spec_end != 0 && self.spec_start == 0 {
return Err(DecodeErrors::FormatStatic(
"Can't merge DC and AC corrupt jpeg"
));
}
// non interleaved data, process one block at a time in trivial scanline order
let k = self.z_order[0];
if k >= self.components.len() {
return Err(DecodeErrors::Format(format!(
"Cannot find component {k}, corrupt image"
)));
}
// For non-interleaved scans, iterate over the component's actual data-unit grid.
let component = &self.components[k];
let mcu_width = (self.info.width as usize * component.horizontal_sample).div_ceil(self.h_max * 8);
let mcu_height = (self.info.height as usize * component.vertical_sample).div_ceil(self.v_max * 8);
for i in 0..mcu_height {
for j in 0..mcu_width {
if self.spec_start != 0 && self.succ_high == 0 && stream.eob_run > 0 {
// handle EOB runs here.
stream.eob_run -= 1;
} else {
let start = 64 * (j + i * (self.components[k].width_stride / 8));
let data: &mut [i16; 64] = buffer
.get_mut(k)
.unwrap()
.get_mut(start..start + 64)
.ok_or(DecodeErrors::FormatStatic("Slice to Small"))?
.try_into()
.unwrap();
if self.spec_start == 0 {
let pos = self.components[k].dc_huff_table & (MAX_COMPONENTS - 1);
let dc_table = self
.dc_huffman_tables
.get(pos)
.ok_or(DecodeErrors::FormatStatic(
"No huffman table for DC component"
))?
.as_ref()
.ok_or(DecodeErrors::FormatStatic(
"Huffman table at index {} not initialized"
))?;
let dc_pred = &mut self.components[k].dc_pred;
if self.succ_high == 0 {
// first scan for this mcu
stream.decode_prog_dc_first(
&mut self.stream,
dc_table,
&mut data[0],
dc_pred
)?;
} else {
// refining scans for this MCU
stream.decode_prog_dc_refine(&mut self.stream, &mut data[0])?;
}
} else {
let pos = self.components[k].ac_huff_table;
let ac_table = self
.ac_huffman_tables
.get(pos)
.ok_or_else(|| {
DecodeErrors::Format(format!(
"No huffman table for component:{pos}"
))
})?
.as_ref()
.ok_or_else(|| {
DecodeErrors::Format(format!(
"Huffman table at index {pos} not initialized"
))
})?;
if self.succ_high == 0 {
debug_assert!(stream.eob_run == 0, "EOB run is not zero");
stream.decode_mcu_ac_first(&mut self.stream, ac_table, data)?;
} else {
// refinement scan
stream.decode_mcu_ac_refine(&mut self.stream, ac_table, data)?;
}
// Check for a marker.
// It can appear in stream CC https://github.com/etemesi254/zune-image/issues/300
// if let Some(marker) = stream.marker.take() {
// self.parse_marker_inner(marker)?;
// }
}
}
// + EOB and investigate effect.
self.todo -= 1;
self.handle_rst_main(stream)?;
}
}
} else {
if self.spec_end != 0 {
return Err(DecodeErrors::HuffmanDecode(
"Can't merge dc and AC corrupt jpeg".to_string()
));
}
// process scan n elements in order
// Do the error checking with allocs here.
// Make the one in the inner loop free of allocations.
for k in 0..self.num_scans {
let n = self.z_order[k as usize];
if n >= self.components.len() {
return Err(DecodeErrors::Format(format!(
"Cannot find component {n}, corrupt image"
)));
}
let component = &mut self.components[n];
let _ = self
.dc_huffman_tables
.get(component.dc_huff_table)
.ok_or_else(|| {
DecodeErrors::Format(format!(
"No huffman table for component:{}",
component.dc_huff_table
))
})?
.as_ref()
.ok_or_else(|| {
DecodeErrors::Format(format!(
"Huffman table at index {} not initialized",
component.dc_huff_table
))
})?;
}
// Interleaved scan
// Components shall not be interleaved in progressive mode, except for
// the DC coefficients in the first scan for each component of a progressive frame.
for i in 0..self.mcu_y {
for j in 0..self.mcu_x {
// process scan n elements in order
for k in 0..self.num_scans {
let n = self.z_order[k as usize];
let component = &mut self.components[n];
let huff_table = self
.dc_huffman_tables
.get(component.dc_huff_table)
.ok_or(DecodeErrors::FormatStatic("No huffman table for component"))?
.as_ref()
.ok_or(DecodeErrors::FormatStatic(
"Huffman table at index not initialized"
))?;
for v_samp in 0..component.vertical_sample {
for h_samp in 0..component.horizontal_sample {
let x2 = j * component.horizontal_sample + h_samp;
let y2 = i * component.vertical_sample + v_samp;
let position = 64 * (x2 + y2 * component.width_stride / 8);
let buf_n = &mut buffer[n];
let Some(data) = &mut buf_n.get_mut(position) else {
// TODO: (CAE), this is another weird sub-sampling bug, so on fix
// remove this
return Err(DecodeErrors::FormatStatic("Invalid image"));
};
if self.succ_high == 0 {
stream.decode_prog_dc_first(
&mut self.stream,
huff_table,
data,
&mut component.dc_pred
)?;
} else {
stream.decode_prog_dc_refine(&mut self.stream, data)?;
}
}
}
}
// We want wrapping subtraction here because it means
// we get a higher number in the case this underflows
self.todo -= 1;
// after every scan that's a mcu, count down restart markers.
self.handle_rst_main(stream)?;
}
}
}
return Ok(());
}
pub(crate) fn handle_rst_main(&mut self, stream: &mut BitStream) -> Result<(), DecodeErrors> {
if self.todo == 0 {
stream.refill(&mut self.stream)?;
}
if self.todo == 0
&& self.restart_interval != 0
&& stream.marker.is_none()
&& !stream.seen_eoi
{
// if no marker and we are to reset RST, look for the marker, this matches
// libjpeg-turbo behaviour and allows us to decode images in
// https://github.com/etemesi254/zune-image/issues/261
let _start = self.stream.position()?;
// skip bytes until we find marker
let marker = get_marker(&mut self.stream, stream);
// In some images, the RST marker on the last section may not be available
// as it is maybe stopped by an EOI marker, see in the case of https://github.com/etemesi254/zune-image/issues/292
// what happened was that we would go looking for the RST marker exhausting all the data
// in the image and this would return an error, so for now
// translate it to a warning, but return the image decoded up
// until that point
if let Ok(marker) = marker {
let _end = self.stream.position()?;
stream.marker = Some(marker);
// NB some warnings may be false positives.
warn!(
"{} Extraneous bytes before marker {:?}",
_end - _start,
marker
);
} else {
warn!("RST marker was not found, where expected, image may be garbled")
}
}
if self.todo == 0 {
self.handle_rst(stream)?
}
Ok(())
}
#[allow(clippy::too_many_lines)]
#[allow(clippy::needless_range_loop, clippy::cast_sign_loss)]
fn finish_progressive_decoding(
&mut self, block: &[Vec<i16>; MAX_COMPONENTS], pixels: &mut [u8]
) -> Result<(), DecodeErrors> {
// This function is complicated because we need to replicate
// the function in mcu.rs
//
// The advantage is that we do very little allocation and very lot
// channel reusing.
// The trick is to notice that we repeat the same procedure per MCU
// width.
//
// So we can set it up that we only allocate temporary storage large enough
// to store a single mcu width, then reuse it per invocation.
//
// This is advantageous to us.
//
// Remember we need to have the whole MCU buffer so we store 3 unprocessed
// channels in memory, and then we allocate the whole output buffer in memory, both of
// which are huge.
//
//
let mcu_height = if self.is_interleaved {
self.mcu_y
} else {
// For non-interleaved images( (1*1) subsampling)
// number of MCU's are the widths (+7 to account for paddings) divided by 8.
self.info.height.div_ceil(8) as usize
};
// Size of our output image(width*height)
let is_hv = usize::from(self.is_interleaved);
let upsampler_scratch_size = is_hv * self.components[0].width_stride;
let width = usize::from(self.info.width);
let padded_width = calculate_padded_width(width, self.info.sample_ratio);
let mut upsampler_scratch_space = vec![0; upsampler_scratch_size];
let mut tmp = [0_i32; DCT_BLOCK];
for (pos, comp) in self.components.iter_mut().enumerate() {
// Allocate only needed components.
//
// For special colorspaces i.e YCCK and CMYK, just allocate all of the needed
// components.
if min(
self.options.jpeg_get_out_colorspace().num_components() - 1,
pos
) == pos
|| self.input_colorspace == ColorSpace::YCCK
|| self.input_colorspace == ColorSpace::CMYK
{
// allocate enough space to hold a whole MCU width
// this means we should take into account sampling ratios
// `*8` is because each MCU spans 8 widths.
let len = comp.width_stride * comp.vertical_sample * 8;
comp.needed = true;
comp.raw_coeff = vec![0; len];
} else {
comp.needed = false;
}
}
let mut pixels_written = 0;
// dequantize, idct and color convert.
for i in 0..mcu_height {
'component: for (position, component) in &mut self.components.iter_mut().enumerate() {
if !component.needed {
continue 'component;
}
let qt_table = &component.quantization_table;
// step is the number of pixels this iteration wil be handling
// Given by the number of mcu's height and the length of the component block
// Since the component block contains the whole channel as raw pixels
// we this evenly divides the pixels into MCU blocks
//
// For interleaved images, this gives us the exact pixels comprising a whole MCU
// block
let step = block[position].len() / mcu_height;
// where we will be reading our pixels from.
let start = i * step;
let slice = &block[position][start..start + step];
let temp_channel = &mut component.raw_coeff;
// The next logical step is to iterate width wise.
// To figure out how many pixels we iterate by we use effective pixels
// Given to us by component.x
// iterate per effective pixels.
let mcu_x = component.width_stride / 8;
// iterate per every vertical sample.
for k in 0..component.vertical_sample {
for j in 0..mcu_x {
// after writing a single stride, we need to skip 8 rows.
// This does the row calculation
let width_stride = k * 8 * component.width_stride;
let start = j * 64 + width_stride;
// See https://github.com/etemesi254/zune-image/issues/262 sample 3.
let Some(qt_slice) = slice.get(start..start + 64) else {
return Err(DecodeErrors::FormatStatic(
"Invalid slice , would panic, invalid image"
));
};
// dequantize
for ((x, out), qt_val) in
qt_slice.iter().zip(tmp.iter_mut()).zip(qt_table.iter())
{
*out = i32::from(*x) * qt_val;
}
// determine where to write.
let sl = &mut temp_channel[component.idct_pos..];
component.idct_pos += 8;
// tmp now contains a dequantized block so idct it
(self.idct_func)(&mut tmp, sl, component.width_stride);
}
// after every write of 8, skip 7 since idct write stride wise 8 times.
//
// Remember each MCU is 8x8 block, so each idct will write 8 strides into
// sl
//
// and component.idct_pos is one stride long
component.idct_pos += 7 * component.width_stride;
}
component.idct_pos = 0;
}
// process that width up until it's impossible
self.post_process(
pixels,
i,
mcu_height,
width,
padded_width,
&mut pixels_written,
&mut upsampler_scratch_space
)?;
}
debug!("Finished decoding image");
return Ok(());
}
pub(crate) fn reset_params(&mut self) {
/*
Apparently, grayscale images which can be down sampled exists, which is weird in the sense
that it has one component Y, which is not usually down sampled.
This means some calculations will be wrong, so for that we explicitly reset params
for such occurrences, warn and reset the image info to appear as if it were
a non-sampled image to ensure decoding works
*/
self.h_max = 1;
self.v_max = 1;
self.info.sample_ratio = SampleRatios::None;
self.is_interleaved = false;
self.components[0].vertical_sample = 1;
self.components[0].width_stride = (((self.info.width as usize) + 7) / 8) * 8;
self.components[0].horizontal_sample = 1;
}
}
///Get a marker from the bit-stream.
///
/// This reads until it gets a marker or end of file is encountered
pub fn get_marker<T>(
reader: &mut ZReader<T>, stream: &mut BitStream
) -> Result<Marker, DecodeErrors>
where
T: ZByteReaderTrait
{
if let Some(marker) = stream.marker {
stream.marker = None;
return Ok(marker);
}
// read until we get a marker
while !reader.eof()? {
let marker = reader.read_u8_err()?;
if marker == 255 {
let mut r = reader.read_u8_err()?;
// 0xFF 0XFF(some images may be like that)
while r == 0xFF {
r = reader.read_u8_err()?;
}
if r != 0 {
return Marker::from_u8(r)
.ok_or_else(|| DecodeErrors::Format(format!("Unknown marker 0xFF{r:X}")));
}
}
}
return Err(DecodeErrors::ExhaustedData);
}
#[cfg(test)]
mod tests{
use zune_core::bytestream::ZCursor;
use crate::JpegDecoder;
#[test]
fn make_test(){
let img = "/Users/etemesi/Downloads/wrong_sampling.jpeg";
let data = ZCursor::new([255, 216, 255, 224, 0, 16, 74, 70, 73, 70, 0, 1, 0, 2, 0, 28, 0, 28, 0, 0, 255, 219, 0, 67, 0, 40, 28, 30, 20, 30, 25, 40, 35, 33, 35, 45, 43, 40, 48, 60, 100, 65, 60, 55, 55, 60, 123, 88, 93, 65, 100, 145, 128, 153, 150, 143, 128, 140, 138, 160, 180, 230, 195, 160, 170, 218, 173, 138, 140, 200, 255, 203, 218, 255, 238, 245, 255, 101, 0, 62, 8, 255, 255, 250, 255, 230, 253, 255, 17, 255, 219, 0, 67, 1, 43, 45, 45, 42, 60, 48, 60, 118, 65, 65, 118, 248, 165, 140, 165, 248, 248, 248, 248, 248, 248, 248, 248, 248, 248, 248, 248, 248, 248, 248, 248, 248, 248, 248, 248, 248, 248, 248, 248, 248, 248, 248, 241, 255, 255, 255, 248, 248, 248, 248, 248, 248, 248, 248, 248, 248, 248, 248, 248, 248, 248, 248, 248, 248, 248, 255, 192, 0, 17, 8, 0, 32, 0, 32, 3, 2, 17, 0, 1, 34, 1, 3, 17, 1, 255, 196, 0, 24, 0, 1, 1, 0, 3, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 5, 3, 0, 1, 4, 255, 196, 0, 37, 16, 0, 2, 2, 1, 4, 1, 3, 5, 0, 0, 0, 0, 0, 0, 0, 0, 1, 2, 3, 17, 0, 4, 18, 33, 48, 34, 65, 81, 113, 19, 20, 51, 97, 161, 255, 196, 0, 22, 1, 1, 1, 1, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 1, 2, 255, 196, 0, 26, 17, 1, 0, 2, 3, 1, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 1, 2, 17, 18, 38, 65, 255, 218, 0, 12, 3, 1, 0, 2, 17, 3, 17, 0, 63, 0, 175, 119, 49, 197, 184, 2, 0, 0, 0, 16, 13, 129, 103, 161, 102, 178, 115, 125, 202, 68, 236, 173, 25, 42, 164, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 4, 0, 38, 0, 0, 0, 0, 250, 255, 255, 255, 0, 0, 0, 0, 0, 0, 0, 67, 1, 43, 45, 45, 60, 48, 60, 118, 65, 65, 118, 248, 165, 140, 165, 248, 248, 248, 248, 248, 248, 248, 248, 248, 248, 248, 248, 248, 248, 248, 248, 248, 248, 248, 248, 248, 248, 248, 248, 248, 248, 248, 241, 255, 255, 255, 248, 248, 248, 248, 248, 248, 248, 248, 248, 248, 248, 248, 248, 248, 248, 248, 248, 248, 248, 255, 192, 0, 17, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 255, 255, 255, 255, 255, 255, 255, 255, 255, 255, 255, 255, 248, 248, 248, 248, 248, 248, 248, 248, 248, 248, 248, 248, 248, 248, 248, 248, 248, 248, 248, 248, 248, 248, 248, 248, 248, 248, 255, 192, 0, 17, 8, 0, 32, 0, 32, 3, 1, 34, 0, 2, 17, 1, 3, 17, 1, 255, 196, 0, 24, 0, 1, 1, 0, 3, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 2, 0, 126, 0, 0, 0, 0, 0, 0, 0, 255, 255, 255, 255, 255, 255, 255, 198]);
let mut decoder = JpegDecoder::new(data);
decoder.decode().unwrap();
}
}
+485
View File
@@ -0,0 +1,485 @@
/*
* Copyright (c) 2023.
*
* This software is free software;
*
* You can redistribute it or modify it under terms of the MIT, Apache License or Zlib license
*/
//!Miscellaneous stuff
#![allow(dead_code)]
use alloc::format;
use core::cmp::max;
use core::fmt;
use core::num::NonZeroU32;
use zune_core::bytestream::ZByteReaderTrait;
use zune_core::colorspace::ColorSpace;
use zune_core::log::{trace, warn};
use crate::components::{ComponentID, SampleRatios};
use crate::errors::DecodeErrors;
use crate::huffman::HuffmanTable;
use crate::JpegDecoder;
/// Start of baseline DCT Huffman coding
pub const START_OF_FRAME_BASE: u16 = 0xffc0;
/// Start of another frame
pub const START_OF_FRAME_EXT_SEQ: u16 = 0xffc1;
/// Start of progressive DCT encoding
pub const START_OF_FRAME_PROG_DCT: u16 = 0xffc2;
/// Start of Lossless sequential Huffman coding
pub const START_OF_FRAME_LOS_SEQ: u16 = 0xffc3;
/// Start of extended sequential DCT arithmetic coding
pub const START_OF_FRAME_EXT_AR: u16 = 0xffc9;
/// Start of Progressive DCT arithmetic coding
pub const START_OF_FRAME_PROG_DCT_AR: u16 = 0xffca;
/// Start of Lossless sequential Arithmetic coding
pub const START_OF_FRAME_LOS_SEQ_AR: u16 = 0xffcb;
/// Undo run length encoding of coefficients by placing them in natural order
///
/// This is an index from position-in-bitstream to position-in-row-major-order.
#[rustfmt::skip]
pub const UN_ZIGZAG: [usize; 64 + 16] = [
0, 1, 8, 16, 9, 2, 3, 10,
17, 24, 32, 25, 18, 11, 4, 5,
12, 19, 26, 33, 40, 48, 41, 34,
27, 20, 13, 6, 7, 14, 21, 28,
35, 42, 49, 56, 57, 50, 43, 36,
29, 22, 15, 23, 30, 37, 44, 51,
58, 59, 52, 45, 38, 31, 39, 46,
53, 60, 61, 54, 47, 55, 62, 63,
// Prevent overflowing
63, 63, 63, 63, 63, 63, 63, 63,
63, 63, 63, 63, 63, 63, 63, 63
];
/// Align data to a 16 byte boundary
#[repr(align(16))]
#[derive(Clone)]
pub struct Aligned16<T: ?Sized>(pub T);
impl<T> Default for Aligned16<T>
where
T: Default
{
fn default() -> Self {
Aligned16(T::default())
}
}
/// Align data to a 32 byte boundary
#[repr(align(32))]
#[derive(Clone)]
pub struct Aligned32<T: ?Sized>(pub T);
impl<T> Default for Aligned32<T>
where
T: Default
{
fn default() -> Self {
Aligned32(T::default())
}
}
/// Markers that identify different Start of Image markers
/// They identify the type of encoding and whether the file use lossy(DCT) or
/// lossless compression and whether we use Huffman or arithmetic coding schemes
#[derive(Eq, PartialEq, Copy, Clone)]
#[allow(clippy::upper_case_acronyms)]
pub enum SOFMarkers {
/// Baseline DCT markers
BaselineDct,
/// SOF_1 Extended sequential DCT,Huffman coding
ExtendedSequentialHuffman,
/// Progressive DCT, Huffman coding
ProgressiveDctHuffman,
/// Lossless (sequential), huffman coding,
LosslessHuffman,
/// Extended sequential DEC, arithmetic coding
ExtendedSequentialDctArithmetic,
/// Progressive DCT, arithmetic coding,
ProgressiveDctArithmetic,
/// Lossless ( sequential), arithmetic coding
LosslessArithmetic
}
impl Default for SOFMarkers {
fn default() -> Self {
Self::BaselineDct
}
}
impl SOFMarkers {
/// Check if a certain marker is sequential DCT or not
pub fn is_sequential_dct(self) -> bool {
matches!(
self,
Self::BaselineDct
| Self::ExtendedSequentialHuffman
| Self::ExtendedSequentialDctArithmetic
)
}
/// Check if a marker is a Lossles type or not
pub fn is_lossless(self) -> bool {
matches!(self, Self::LosslessHuffman | Self::LosslessArithmetic)
}
/// Check whether a marker is a progressive marker or not
pub fn is_progressive(self) -> bool {
matches!(
self,
Self::ProgressiveDctHuffman | Self::ProgressiveDctArithmetic
)
}
/// Create a marker from an integer
pub fn from_int(int: u16) -> Option<SOFMarkers> {
match int {
START_OF_FRAME_BASE => Some(Self::BaselineDct),
START_OF_FRAME_PROG_DCT => Some(Self::ProgressiveDctHuffman),
START_OF_FRAME_PROG_DCT_AR => Some(Self::ProgressiveDctArithmetic),
START_OF_FRAME_LOS_SEQ => Some(Self::LosslessHuffman),
START_OF_FRAME_LOS_SEQ_AR => Some(Self::LosslessArithmetic),
START_OF_FRAME_EXT_SEQ => Some(Self::ExtendedSequentialHuffman),
START_OF_FRAME_EXT_AR => Some(Self::ExtendedSequentialDctArithmetic),
_ => None
}
}
}
impl fmt::Debug for SOFMarkers {
fn fmt(&self, f: &mut fmt::Formatter) -> fmt::Result {
match &self {
Self::BaselineDct => write!(f, "Baseline DCT"),
Self::ExtendedSequentialHuffman => {
write!(f, "Extended sequential DCT, Huffman Coding")
}
Self::ProgressiveDctHuffman => write!(f, "Progressive DCT,Huffman Encoding"),
Self::LosslessHuffman => write!(f, "Lossless (sequential) Huffman encoding"),
Self::ExtendedSequentialDctArithmetic => {
write!(f, "Extended sequential DCT, arithmetic coding")
}
Self::ProgressiveDctArithmetic => write!(f, "Progressive DCT, arithmetic coding"),
Self::LosslessArithmetic => write!(f, "Lossless (sequential) arithmetic coding")
}
}
}
/// Set up component parameters.
///
/// This modifies the components in place setting up details needed by other
/// parts fo the decoder.
pub(crate) fn setup_component_params<T: ZByteReaderTrait>(
img: &mut JpegDecoder<T>
) -> Result<(), DecodeErrors> {
let img_width = img.width();
let img_height = img.height();
// in case of adobe app14 being present, zero may indicate
// either CMYK if components are 4 or RGB if components are 3,
// see https://docs.oracle.com/javase/6/docs/api/javax/imageio/metadata/doc-files/jpeg_metadata.html
// so since we may not know how many number of components
// we have when decoding app14, we have to defer that check
// until now.
//
// We know adobe app14 was present since it's the only one that can modify
// input colorspace to be CMYK
if img.components.len() == 3 && img.input_colorspace == ColorSpace::CMYK {
img.input_colorspace = ColorSpace::RGB;
}
for component in &mut img.components {
// compute interleaved image info
// h_max contains the maximum horizontal component
img.h_max = max(img.h_max, component.horizontal_sample);
// v_max contains the maximum vertical component
img.v_max = max(img.v_max, component.vertical_sample);
img.mcu_width = img.h_max * 8;
img.mcu_height = img.v_max * 8;
// Number of MCU's per width
img.mcu_x = usize::from(img.info.width).div_ceil(img.mcu_width);
// Number of MCU's per height
img.mcu_y = usize::from(img.info.height).div_ceil(img.mcu_height);
if img.h_max != 1 || img.v_max != 1 {
// interleaved images have horizontal and vertical sampling factors
// not equal to 1.
img.is_interleaved = true;
}
// Extract quantization tables from the arrays into components
let qt_table = *img.qt_tables[component.quantization_table_number as usize]
.as_ref()
.ok_or_else(|| {
DecodeErrors::DqtError(format!(
"No quantization table for component {:?}",
component.component_id
))
})?;
let x = (usize::from(img_width) * component.horizontal_sample + img.h_max - 1) / img.h_max;
let y = (usize::from(img_height) * component.horizontal_sample + img.h_max - 1) / img.v_max;
component.x = x;
component.w2 = img.mcu_x * component.horizontal_sample * 8;
// probably not needed. :)
component.y = y;
component.quantization_table = qt_table;
// initially stride contains its horizontal sub-sampling
component.width_stride *= img.mcu_x * 8;
}
{
// Sampling factors are one thing that suck
// this fixes a specific problem with images like
//
// (2 2) None
// (2 1) H
// (2 1) H
//
// The images exist in the wild, the images are not meant to exist
// but they do, it's just an annoying horizontal sub-sampling that
// I don't know why it exists.
// But it does
// So we try to cope with that.
// I am not sure of how to explain how to fix it, but it involved a debugger
// and to much coke(the legal one)
//
// If this wasn't present, self.upsample_dest would have the wrong length
let mut handle_that_annoying_bug = false;
if let Some(y_component) = img
.components
.iter()
.find(|c| c.component_id == ComponentID::Y)
{
if y_component.horizontal_sample == 2 || y_component.vertical_sample == 2 {
handle_that_annoying_bug = true;
}
}
if handle_that_annoying_bug {
for comp in &mut img.components {
if (comp.component_id != ComponentID::Y)
&& (comp.horizontal_sample != 1 || comp.vertical_sample != 1)
{
comp.fix_an_annoying_bug = 2;
}
}
}
}
if img.is_mjpeg {
fill_default_mjpeg_tables(
img.is_progressive,
&mut img.dc_huffman_tables,
&mut img.ac_huffman_tables
);
}
// check colorspace matches
if img.input_colorspace.num_components() > img.components.len() {
if img.input_colorspace == ColorSpace::YCCK {
// Some images may have YCCK format (from adobe app14 segment) which is supposed to be 4 components
// but only 3 components, see issue https://github.com/etemesi254/zune-image/issues/275
// So this is the behaviour of other decoders
// - stb_image: Treats it as YCbCr image
// - libjpeg_turbo: Does not know how to parse YCCK images (transform 2 app14) so treats
// it as YCbCr
// So I will match that to match existing ones
warn!("Treating YCCK colorspace as YCbCr as component length does not match");
img.input_colorspace = ColorSpace::YCbCr
} else {
// Note, translated this to a warning to handle valid images of the sort
// See https://github.com/etemesi254/zune-image/issues/288 where there
// was a CMYK image with two components which would be decoded to 4 components
// by the decoder.
// So with a warning that becomes supported.
//
// djpeg fails to render an image from that also probably because it does not
// understand the expected format.
if !img.options.strict_mode() {
warn!(
"Expected {} number of components but found {}",
img.input_colorspace.num_components(),
img.components.len()
);
warn!("Defaulting to multisample to decode");
// N/B: We do not post process the color of such, treating it as multiband
// is the best option since I am not aware of grayscale+alpha which is the most common
// two band format in jpeg.
if img.components.len() > 0 {
img.input_colorspace = ColorSpace::MultiBand(
NonZeroU32::new(img.components.len() as u32).unwrap()
);
}
} else {
let msg = format!(
"Expected {} number of components but found {}",
img.input_colorspace.num_components(),
img.components.len()
);
return Err(DecodeErrors::Format(msg));
}
}
}
Ok(())
}
///Calculate number of fill bytes added to the end of a JPEG image
/// to fill the image
///
/// JPEG usually inserts padding bytes if the image width cannot be evenly divided into
/// 8 , 16 or 32 chunks depending on the sub sampling ratio. So given a sub-sampling ratio,
/// and the actual width, this calculates the padded bytes that were added to the image
///
/// # Params
/// -actual_width: Actual width of the image
/// -sub_sample: Sub sampling factor of the image
///
/// # Returns
/// The padded width, this is how long the width is for a particular image
pub fn calculate_padded_width(actual_width: usize, sub_sample: SampleRatios) -> usize {
match sub_sample {
SampleRatios::None | SampleRatios::V => {
// None+V sends one MCU row, so that's a simple calculation
((actual_width + 7) / 8) * 8
}
SampleRatios::H | SampleRatios::HV => {
// sends two rows, width can be expanded by up to 15 more bytes
((actual_width + 15) / 16) * 16
}
SampleRatios::Generic(h, _) => {
((actual_width + ((h * 8).saturating_sub(1))) / (h * 8)) * (h * 8)
}
}
}
// https://www.loc.gov/preservation/digital/formats/fdd/fdd000063.shtml
// "Avery Lee, writing in the rec.video.desktop newsgroup in 2001, commented that "MJPEG, or at
// least the MJPEG in AVIs having the MJPG fourcc, is restricted JPEG with a fixed -- and
// *omitted* -- Huffman table. The JPEG must be YCbCr colorspace, it must be 4:2:2, and it must
// use basic Huffman encoding, not arithmetic or progressive.... You can indeed extract the
// MJPEG frames and decode them with a regular JPEG decoder, but you have to prepend the DHT
// segment to them, or else the decoder won't have any idea how to decompress the data.
// The exact table necessary is given in the OpenDML spec.""
pub fn fill_default_mjpeg_tables(
is_progressive: bool, dc_huffman_tables: &mut [Option<HuffmanTable>],
ac_huffman_tables: &mut [Option<HuffmanTable>]
) {
// Section K.3.3
trace!("Filling with default mjpeg tables");
if dc_huffman_tables[0].is_none() {
// Table K.3
dc_huffman_tables[0] = Some(
HuffmanTable::new_unfilled(
&[
0x00, 0x00, 0x01, 0x05, 0x01, 0x01, 0x01, 0x01, 0x01, 0x01, 0x00, 0x00, 0x00,
0x00, 0x00, 0x00, 0x00
],
&[
0x00, 0x01, 0x02, 0x03, 0x04, 0x05, 0x06, 0x07, 0x08, 0x09, 0x0A, 0x0B
],
true,
is_progressive
)
.unwrap()
);
}
if dc_huffman_tables[1].is_none() {
// Table K.4
dc_huffman_tables[1] = Some(
HuffmanTable::new_unfilled(
&[
0x00, 0x00, 0x03, 0x01, 0x01, 0x01, 0x01, 0x01, 0x01, 0x01, 0x01, 0x01, 0x00,
0x00, 0x00, 0x00, 0x00
],
&[
0x00, 0x01, 0x02, 0x03, 0x04, 0x05, 0x06, 0x07, 0x08, 0x09, 0x0A, 0x0B
],
true,
is_progressive
)
.unwrap()
);
}
if ac_huffman_tables[0].is_none() {
// Table K.5
ac_huffman_tables[0] = Some(
HuffmanTable::new_unfilled(
&[
0x00, 0x00, 0x02, 0x01, 0x03, 0x03, 0x02, 0x04, 0x03, 0x05, 0x05, 0x04, 0x04,
0x00, 0x00, 0x01, 0x7D
],
&[
0x01, 0x02, 0x03, 0x00, 0x04, 0x11, 0x05, 0x12, 0x21, 0x31, 0x41, 0x06, 0x13,
0x51, 0x61, 0x07, 0x22, 0x71, 0x14, 0x32, 0x81, 0x91, 0xA1, 0x08, 0x23, 0x42,
0xB1, 0xC1, 0x15, 0x52, 0xD1, 0xF0, 0x24, 0x33, 0x62, 0x72, 0x82, 0x09, 0x0A,
0x16, 0x17, 0x18, 0x19, 0x1A, 0x25, 0x26, 0x27, 0x28, 0x29, 0x2A, 0x34, 0x35,
0x36, 0x37, 0x38, 0x39, 0x3A, 0x43, 0x44, 0x45, 0x46, 0x47, 0x48, 0x49, 0x4A,
0x53, 0x54, 0x55, 0x56, 0x57, 0x58, 0x59, 0x5A, 0x63, 0x64, 0x65, 0x66, 0x67,
0x68, 0x69, 0x6A, 0x73, 0x74, 0x75, 0x76, 0x77, 0x78, 0x79, 0x7A, 0x83, 0x84,
0x85, 0x86, 0x87, 0x88, 0x89, 0x8A, 0x92, 0x93, 0x94, 0x95, 0x96, 0x97, 0x98,
0x99, 0x9A, 0xA2, 0xA3, 0xA4, 0xA5, 0xA6, 0xA7, 0xA8, 0xA9, 0xAA, 0xB2, 0xB3,
0xB4, 0xB5, 0xB6, 0xB7, 0xB8, 0xB9, 0xBA, 0xC2, 0xC3, 0xC4, 0xC5, 0xC6, 0xC7,
0xC8, 0xC9, 0xCA, 0xD2, 0xD3, 0xD4, 0xD5, 0xD6, 0xD7, 0xD8, 0xD9, 0xDA, 0xE1,
0xE2, 0xE3, 0xE4, 0xE5, 0xE6, 0xE7, 0xE8, 0xE9, 0xEA, 0xF1, 0xF2, 0xF3, 0xF4,
0xF5, 0xF6, 0xF7, 0xF8, 0xF9, 0xFA
],
false,
is_progressive
)
.unwrap()
);
}
if ac_huffman_tables[1].is_none() {
// Table K.6
ac_huffman_tables[1] = Some(
HuffmanTable::new_unfilled(
&[
0x00, 0x00, 0x02, 0x01, 0x02, 0x04, 0x04, 0x03, 0x04, 0x07, 0x05, 0x04, 0x04,
0x00, 0x01, 0x02, 0x77
],
&[
0x00, 0x01, 0x02, 0x03, 0x11, 0x04, 0x05, 0x21, 0x31, 0x06, 0x12, 0x41, 0x51,
0x07, 0x61, 0x71, 0x13, 0x22, 0x32, 0x81, 0x08, 0x14, 0x42, 0x91, 0xA1, 0xB1,
0xC1, 0x09, 0x23, 0x33, 0x52, 0xF0, 0x15, 0x62, 0x72, 0xD1, 0x0A, 0x16, 0x24,
0x34, 0xE1, 0x25, 0xF1, 0x17, 0x18, 0x19, 0x1A, 0x26, 0x27, 0x28, 0x29, 0x2A,
0x35, 0x36, 0x37, 0x38, 0x39, 0x3A, 0x43, 0x44, 0x45, 0x46, 0x47, 0x48, 0x49,
0x4A, 0x53, 0x54, 0x55, 0x56, 0x57, 0x58, 0x59, 0x5A, 0x63, 0x64, 0x65, 0x66,
0x67, 0x68, 0x69, 0x6A, 0x73, 0x74, 0x75, 0x76, 0x77, 0x78, 0x79, 0x7A, 0x82,
0x83, 0x84, 0x85, 0x86, 0x87, 0x88, 0x89, 0x8A, 0x92, 0x93, 0x94, 0x95, 0x96,
0x97, 0x98, 0x99, 0x9A, 0xA2, 0xA3, 0xA4, 0xA5, 0xA6, 0xA7, 0xA8, 0xA9, 0xAA,
0xB2, 0xB3, 0xB4, 0xB5, 0xB6, 0xB7, 0xB8, 0xB9, 0xBA, 0xC2, 0xC3, 0xC4, 0xC5,
0xC6, 0xC7, 0xC8, 0xC9, 0xCA, 0xD2, 0xD3, 0xD4, 0xD5, 0xD6, 0xD7, 0xD8, 0xD9,
0xDA, 0xE2, 0xE3, 0xE4, 0xE5, 0xE6, 0xE7, 0xE8, 0xE9, 0xEA, 0xF2, 0xF3, 0xF4,
0xF5, 0xF6, 0xF7, 0xF8, 0xF9, 0xFA
],
false,
is_progressive
)
.unwrap()
);
}
}
+4
View File
@@ -0,0 +1,4 @@
#[cfg(all(feature = "x86", any(target_arch = "x86", target_arch = "x86_64")))]
pub use crate::unsafe_utils_avx2::*;
#[cfg(all(feature = "neon", target_arch = "aarch64"))]
pub use crate::unsafe_utils_neon::*;
+223
View File
@@ -0,0 +1,223 @@
/*
* Copyright (c) 2023.
*
* This software is free software;
*
* You can redistribute it or modify it under terms of the MIT, Apache License or Zlib license
*/
#![cfg(all(feature = "x86", any(target_arch = "x86", target_arch = "x86_64")))]
//! This module provides unsafe ways to do some things
#![allow(clippy::wildcard_imports)]
#[cfg(target_arch = "x86")]
use core::arch::x86::*;
#[cfg(target_arch = "x86_64")]
use core::arch::x86_64::*;
use core::ops::{Add, AddAssign, Mul, MulAssign, Sub};
/// A copy of `_MM_SHUFFLE()` that doesn't require
/// a nightly compiler
#[inline]
const fn shuffle(z: i32, y: i32, x: i32, w: i32) -> i32 {
(z << 6) | (y << 4) | (x << 2) | w
}
/// An abstraction of an AVX ymm register that
///allows some things to not look ugly
#[derive(Clone, Copy)]
pub struct YmmRegister {
/// An AVX register
pub(crate) mm256: __m256i
}
impl Add for YmmRegister {
type Output = YmmRegister;
#[inline]
fn add(self, rhs: Self) -> Self::Output {
unsafe {
return YmmRegister {
mm256: _mm256_add_epi32(self.mm256, rhs.mm256)
};
}
}
}
impl Add<i32> for YmmRegister {
type Output = YmmRegister;
#[inline]
fn add(self, rhs: i32) -> Self::Output {
unsafe {
let tmp = _mm256_set1_epi32(rhs);
return YmmRegister {
mm256: _mm256_add_epi32(self.mm256, tmp)
};
}
}
}
impl Sub for YmmRegister {
type Output = YmmRegister;
#[inline]
fn sub(self, rhs: Self) -> Self::Output {
unsafe {
return YmmRegister {
mm256: _mm256_sub_epi32(self.mm256, rhs.mm256)
};
}
}
}
impl AddAssign for YmmRegister {
#[inline]
fn add_assign(&mut self, rhs: Self) {
unsafe {
self.mm256 = _mm256_add_epi32(self.mm256, rhs.mm256);
}
}
}
impl AddAssign<i32> for YmmRegister {
#[inline]
fn add_assign(&mut self, rhs: i32) {
unsafe {
let tmp = _mm256_set1_epi32(rhs);
self.mm256 = _mm256_add_epi32(self.mm256, tmp);
}
}
}
impl Mul for YmmRegister {
type Output = YmmRegister;
#[inline]
fn mul(self, rhs: Self) -> Self::Output {
unsafe {
YmmRegister {
mm256: _mm256_mullo_epi32(self.mm256, rhs.mm256)
}
}
}
}
impl Mul<i32> for YmmRegister {
type Output = YmmRegister;
#[inline]
fn mul(self, rhs: i32) -> Self::Output {
unsafe {
let tmp = _mm256_set1_epi32(rhs);
YmmRegister {
mm256: _mm256_mullo_epi32(self.mm256, tmp)
}
}
}
}
impl MulAssign for YmmRegister {
#[inline]
fn mul_assign(&mut self, rhs: Self) {
unsafe {
self.mm256 = _mm256_mullo_epi32(self.mm256, rhs.mm256);
}
}
}
impl MulAssign<i32> for YmmRegister {
#[inline]
fn mul_assign(&mut self, rhs: i32) {
unsafe {
let tmp = _mm256_set1_epi32(rhs);
self.mm256 = _mm256_mullo_epi32(self.mm256, tmp);
}
}
}
impl MulAssign<__m256i> for YmmRegister {
#[inline]
fn mul_assign(&mut self, rhs: __m256i) {
unsafe {
self.mm256 = _mm256_mullo_epi32(self.mm256, rhs);
}
}
}
type Reg = YmmRegister;
/// Transpose an array of 8 by 8 i32's using avx intrinsics
///
/// This was translated from [here](https://newbedev.com/transpose-an-8x8-float-using-avx-avx2)
#[allow(unused_parens, clippy::too_many_arguments)]
#[target_feature(enable = "avx2")]
#[inline]
pub unsafe fn transpose(
v0: &mut Reg, v1: &mut Reg, v2: &mut Reg, v3: &mut Reg, v4: &mut Reg, v5: &mut Reg,
v6: &mut Reg, v7: &mut Reg
) {
macro_rules! merge_epi32 {
($v0:tt,$v1:tt,$v2:tt,$v3:tt) => {
let va = _mm256_permute4x64_epi64($v0, shuffle(3, 1, 2, 0));
let vb = _mm256_permute4x64_epi64($v1, shuffle(3, 1, 2, 0));
$v2 = _mm256_unpacklo_epi32(va, vb);
$v3 = _mm256_unpackhi_epi32(va, vb);
};
}
macro_rules! merge_epi64 {
($v0:tt,$v1:tt,$v2:tt,$v3:tt) => {
let va = _mm256_permute4x64_epi64($v0, shuffle(3, 1, 2, 0));
let vb = _mm256_permute4x64_epi64($v1, shuffle(3, 1, 2, 0));
$v2 = _mm256_unpacklo_epi64(va, vb);
$v3 = _mm256_unpackhi_epi64(va, vb);
};
}
macro_rules! merge_si128 {
($v0:tt,$v1:tt,$v2:tt,$v3:tt) => {
$v2 = _mm256_permute2x128_si256($v0, $v1, shuffle(0, 2, 0, 0));
$v3 = _mm256_permute2x128_si256($v0, $v1, shuffle(0, 3, 0, 1));
};
}
let (w0, w1, w2, w3, w4, w5, w6, w7);
merge_epi32!((v0.mm256), (v1.mm256), w0, w1);
merge_epi32!((v2.mm256), (v3.mm256), w2, w3);
merge_epi32!((v4.mm256), (v5.mm256), w4, w5);
merge_epi32!((v6.mm256), (v7.mm256), w6, w7);
let (x0, x1, x2, x3, x4, x5, x6, x7);
merge_epi64!(w0, w2, x0, x1);
merge_epi64!(w1, w3, x2, x3);
merge_epi64!(w4, w6, x4, x5);
merge_epi64!(w5, w7, x6, x7);
merge_si128!(x0, x4, (v0.mm256), (v1.mm256));
merge_si128!(x1, x5, (v2.mm256), (v3.mm256));
merge_si128!(x2, x6, (v4.mm256), (v5.mm256));
merge_si128!(x3, x7, (v6.mm256), (v7.mm256));
}
+331
View File
@@ -0,0 +1,331 @@
/*
* Copyright (c) 2023.
*
* This software is free software;
*
* You can redistribute it or modify it under terms of the MIT, Apache License or Zlib license
*/
#![cfg(all(feature = "neon", target_arch = "aarch64"))]
// TODO can this be extended to armv7
//! This module provides unsafe ways to do some things
#![allow(clippy::wildcard_imports)]
use core::arch::aarch64::*;
use core::ops::{Add, AddAssign, BitOr, BitOrAssign, Mul, MulAssign, Sub};
pub type VecType = int32x4x2_t;
pub unsafe fn loadu(src: *const i32) -> VecType {
vld1q_s32_x2(src as *const _)
}
/// An abstraction of an AVX ymm register that
///allows some things to not look ugly
#[derive(Clone, Copy)]
pub struct YmmRegister {
/// An AVX register
pub(crate) mm256: VecType
}
impl YmmRegister {
#[inline]
pub unsafe fn load(src: *const i32) -> Self {
loadu(src).into()
}
#[inline]
pub fn map2(self, other: Self, f: impl Fn(int32x4_t, int32x4_t) -> int32x4_t) -> Self {
let m0 = f(self.mm256.0, other.mm256.0);
let m1 = f(self.mm256.1, other.mm256.1);
YmmRegister {
mm256: int32x4x2_t(m0, m1)
}
}
#[inline]
pub fn all_zero(self) -> bool {
unsafe {
let both = vorrq_s32(self.mm256.0, self.mm256.1);
let both_unsigned = vreinterpretq_u32_s32(both);
0 == vmaxvq_u32(both_unsigned)
}
}
#[inline]
pub fn const_shl<const N: i32>(self) -> Self {
// Ensure that we logically shift left
unsafe {
let m0 = vreinterpretq_s32_u32(vshlq_n_u32::<N>(vreinterpretq_u32_s32(self.mm256.0)));
let m1 = vreinterpretq_s32_u32(vshlq_n_u32::<N>(vreinterpretq_u32_s32(self.mm256.1)));
YmmRegister {
mm256: int32x4x2_t(m0, m1)
}
}
}
#[inline]
pub fn const_shra<const N: i32>(self) -> Self {
unsafe {
let i0 = vshrq_n_s32::<N>(self.mm256.0);
let i1 = vshrq_n_s32::<N>(self.mm256.1);
YmmRegister {
mm256: int32x4x2_t(i0, i1)
}
}
}
}
impl<T> Add<T> for YmmRegister
where
T: Into<Self>
{
type Output = YmmRegister;
#[inline]
fn add(self, rhs: T) -> Self::Output {
let rhs = rhs.into();
unsafe { self.map2(rhs, |a, b| vaddq_s32(a, b)) }
}
}
impl<T> Sub<T> for YmmRegister
where
T: Into<Self>
{
type Output = YmmRegister;
#[inline]
fn sub(self, rhs: T) -> Self::Output {
let rhs = rhs.into();
unsafe { self.map2(rhs, |a, b| vsubq_s32(a, b)) }
}
}
impl<T> AddAssign<T> for YmmRegister
where
T: Into<Self>
{
#[inline]
fn add_assign(&mut self, rhs: T) {
let rhs: Self = rhs.into();
*self = *self + rhs;
}
}
impl<T> Mul<T> for YmmRegister
where
T: Into<Self>
{
type Output = YmmRegister;
#[inline]
fn mul(self, rhs: T) -> Self::Output {
let rhs = rhs.into();
unsafe { self.map2(rhs, |a, b| vmulq_s32(a, b)) }
}
}
impl<T> MulAssign<T> for YmmRegister
where
T: Into<Self>
{
#[inline]
fn mul_assign(&mut self, rhs: T) {
let rhs: Self = rhs.into();
*self = *self * rhs;
}
}
impl<T> BitOr<T> for YmmRegister
where
T: Into<Self>
{
type Output = YmmRegister;
#[inline]
fn bitor(self, rhs: T) -> Self::Output {
let rhs = rhs.into();
unsafe { self.map2(rhs, |a, b| vorrq_s32(a, b)) }
}
}
impl<T> BitOrAssign<T> for YmmRegister
where
T: Into<Self>
{
#[inline]
fn bitor_assign(&mut self, rhs: T) {
let rhs: Self = rhs.into();
*self = *self | rhs;
}
}
impl From<i32> for YmmRegister {
#[inline]
fn from(val: i32) -> Self {
unsafe {
let dup = vdupq_n_s32(val);
YmmRegister {
mm256: int32x4x2_t(dup, dup)
}
}
}
}
impl From<VecType> for YmmRegister {
#[inline]
fn from(mm256: VecType) -> Self {
YmmRegister { mm256 }
}
}
#[allow(clippy::too_many_arguments)]
#[inline]
unsafe fn transpose4(
v0: &mut int32x4_t, v1: &mut int32x4_t, v2: &mut int32x4_t, v3: &mut int32x4_t
) {
let w0 = vtrnq_s32(
vreinterpretq_s32_s64(vtrn1q_s64(
vreinterpretq_s64_s32(*v0),
vreinterpretq_s64_s32(*v2)
)),
vreinterpretq_s32_s64(vtrn1q_s64(
vreinterpretq_s64_s32(*v1),
vreinterpretq_s64_s32(*v3)
))
);
let w1 = vtrnq_s32(
vreinterpretq_s32_s64(vtrn2q_s64(
vreinterpretq_s64_s32(*v0),
vreinterpretq_s64_s32(*v2)
)),
vreinterpretq_s32_s64(vtrn2q_s64(
vreinterpretq_s64_s32(*v1),
vreinterpretq_s64_s32(*v3)
))
);
*v0 = w0.0;
*v1 = w0.1;
*v2 = w1.0;
*v3 = w1.1;
}
/// Transpose an array of 8 by 8 i32
/// Arm has dedicated interleave/transpose instructions
/// we:
/// 1. Transpose the upper left and lower right quadrants
/// 2. Swap and transpose the upper right and lower left quadrants
#[allow(clippy::too_many_arguments)]
#[inline]
pub unsafe fn transpose(
v0: &mut YmmRegister, v1: &mut YmmRegister, v2: &mut YmmRegister, v3: &mut YmmRegister,
v4: &mut YmmRegister, v5: &mut YmmRegister, v6: &mut YmmRegister, v7: &mut YmmRegister
) {
use core::mem::swap;
let ul0 = &mut v0.mm256.0;
let ul1 = &mut v1.mm256.0;
let ul2 = &mut v2.mm256.0;
let ul3 = &mut v3.mm256.0;
let ur0 = &mut v0.mm256.1;
let ur1 = &mut v1.mm256.1;
let ur2 = &mut v2.mm256.1;
let ur3 = &mut v3.mm256.1;
let ll0 = &mut v4.mm256.0;
let ll1 = &mut v5.mm256.0;
let ll2 = &mut v6.mm256.0;
let ll3 = &mut v7.mm256.0;
let lr0 = &mut v4.mm256.1;
let lr1 = &mut v5.mm256.1;
let lr2 = &mut v6.mm256.1;
let lr3 = &mut v7.mm256.1;
swap(ur0, ll0);
swap(ur1, ll1);
swap(ur2, ll2);
swap(ur3, ll3);
transpose4(ul0, ul1, ul2, ul3);
transpose4(ur0, ur1, ur2, ur3);
transpose4(ll0, ll1, ll2, ll3);
transpose4(lr0, lr1, lr2, lr3);
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_transpose() {
fn get_val(i: usize, j: usize) -> i32 {
((i * 8) / (j + 1)) as i32
}
unsafe {
let mut vals: [i32; 8 * 8] = [0; 8 * 8];
for i in 0..8 {
for j in 0..8 {
// some order-dependent value of i and j
let value = get_val(i, j);
vals[i * 8 + j] = value;
}
}
let mut regs: [YmmRegister; 8] = core::mem::transmute(vals);
let mut reg0 = regs[0];
let mut reg1 = regs[1];
let mut reg2 = regs[2];
let mut reg3 = regs[3];
let mut reg4 = regs[4];
let mut reg5 = regs[5];
let mut reg6 = regs[6];
let mut reg7 = regs[7];
transpose(
&mut reg0, &mut reg1, &mut reg2, &mut reg3, &mut reg4, &mut reg5, &mut reg6,
&mut reg7
);
regs[0] = reg0;
regs[1] = reg1;
regs[2] = reg2;
regs[3] = reg3;
regs[4] = reg4;
regs[5] = reg5;
regs[6] = reg6;
regs[7] = reg7;
let vals_from_reg: [i32; 8 * 8] = core::mem::transmute(regs);
for i in 0..8 {
for j in 0..i {
let orig = vals[i * 8 + j];
vals[i * 8 + j] = vals[j * 8 + i];
vals[j * 8 + i] = orig;
}
}
for i in 0..8 {
for j in 0..8 {
assert_eq!(vals[j * 8 + i], get_val(i, j));
assert_eq!(vals_from_reg[j * 8 + i], get_val(i, j));
}
}
assert_eq!(vals, vals_from_reg);
}
}
}
+343
View File
@@ -0,0 +1,343 @@
/*
* Copyright (c) 2023.
*
* This software is free software;
*
* You can redistribute it or modify it under terms of the MIT, Apache License or Zlib license
*/
//! Up-sampling routines
//!
//! The main upsampling method is a bi-linear interpolation or a "triangle
//! filter " or libjpeg turbo `fancy_upsampling` which is a good compromise
//! between speed and visual quality
//!
//! # The filter
//! Each output pixel is made from `(3*A+B)/4` where A is the original
//! pixel closer to the output and B is the one further.
//!
//! ```text
//!+---+---+
//! | A | B |
//! +---+---+
//! +-+-+-+-+
//! | |P| | |
//! +-+-+-+-+
//! ```
//!
//! # Horizontal Bi-linear filter
//! ```text
//! |---+-----------+---+
//! | | | |
//! | A | |p1 | p2| | B |
//! | | | |
//! |---+-----------+---+
//!
//! ```
//! For a horizontal bi-linear it's trivial to implement,
//!
//! `A` becomes the input closest to the output.
//!
//! `B` varies depending on output.
//! - For odd positions, input is the `next` pixel after A
//! - For even positions, input is the `previous` value before A.
//!
//! We iterate in a classic 1-D sliding window with a window of 3.
//! For our sliding window approach, `A` is the 1st and `B` is either the 0th term or 2nd term
//! depending on position we are writing.(see scalar code).
//!
//! For vector code see module sse for explanation.
//!
//! # Vertical bi-linear.
//! Vertical up-sampling is a bit trickier.
//!
//! ```text
//! +----+----+
//! | A1 | A2 |
//! +----+----+
//! +----+----+
//! | p1 | p2 |
//! +----+-+--+
//! +----+-+--+
//! | p3 | p4 |
//! +----+-+--+
//! +----+----+
//! | B1 | B2 |
//! +----+----+
//! ```
//!
//! For `p1`
//! - `A1` is given a weight of `3` and `B1` is given a weight of 1.
//!
//! For `p3`
//! - `B1` is given a weight of `3` and `A1` is given a weight of 1
//!
//! # Horizontal vertical downsampling/chroma quartering.
//!
//! Carry out a vertical filter in the first pass, then a horizontal filter in the second pass.
#![allow(unreachable_code)]
use zune_core::options::DecoderOptions;
use crate::components::UpSampler;
#[cfg(any(target_arch = "x86", target_arch = "x86_64"))]
#[cfg(feature = "x86")]
mod avx2;
#[cfg(target_arch = "aarch64")]
#[cfg(feature = "neon")]
mod neon;
#[cfg(feature = "portable_simd")]
mod portable_simd;
mod scalar;
// choose the best possible implementation for this platform
#[allow(unused_variables)]
pub fn choose_horizontal_samp_function(options: &DecoderOptions) -> UpSampler {
#[cfg(any(target_arch = "x86", target_arch = "x86_64"))]
#[cfg(feature = "x86")]
if options.use_avx2() {
return |a: &[i16], b: &[i16], c: &[i16], d: &mut [i16], e: &mut [i16]| {
// SAFETY: `options.use_avx2()` only returns true if avx2 is supported.
unsafe { avx2::upsample_horizontal_avx2(a, b, c, d, e) }
};
}
#[cfg(target_arch = "aarch64")]
#[cfg(feature = "neon")]
if options.use_neon() {
return |a: &[i16], b: &[i16], c: &[i16], d: &mut [i16], e: &mut [i16]| {
// SAFETY: `options.use_neon()` only returns true if neon is supported.
unsafe { neon::upsample_horizontal_neon(a, b, c, d, e) }
};
}
#[cfg(feature = "portable_simd")]
return portable_simd::upsample_horizontal_simd;
return scalar::upsample_horizontal;
}
#[allow(unused_variables)]
pub fn choose_hv_samp_function(options: &DecoderOptions) -> UpSampler {
#[cfg(any(target_arch = "x86", target_arch = "x86_64"))]
#[cfg(feature = "x86")]
if options.use_avx2() {
return |a: &[i16], b: &[i16], c: &[i16], d: &mut [i16], e: &mut [i16]| {
// SAFETY: `options.use_avx2()` only returns true if avx2 is supported.
unsafe { avx2::upsample_hv_avx2(a, b, c, d, e) }
};
}
#[cfg(target_arch = "aarch64")]
#[cfg(feature = "neon")]
if options.use_neon() {
return |a: &[i16], b: &[i16], c: &[i16], d: &mut [i16], e: &mut [i16]| {
// SAFETY: `options.use_neon()` only returns true if neon is supported.
unsafe { neon::upsample_hv_neon(a, b, c, d, e) }
};
}
#[cfg(feature = "portable_simd")]
return portable_simd::upsample_hv_simd;
return scalar::upsample_hv;
}
#[allow(unused_variables)]
pub fn choose_v_samp_function(options: &DecoderOptions) -> UpSampler {
#[cfg(any(target_arch = "x86", target_arch = "x86_64"))]
#[cfg(feature = "x86")]
if options.use_avx2() {
return |a: &[i16], b: &[i16], c: &[i16], d: &mut [i16], e: &mut [i16]| {
// SAFETY: `options.use_avx2()` only returns true if avx2 is supported.
unsafe { avx2::upsample_vertical_avx2(a, b, c, d, e) }
};
}
#[cfg(target_arch = "aarch64")]
#[cfg(feature = "neon")]
if options.use_neon() {
return |a: &[i16], b: &[i16], c: &[i16], d: &mut [i16], e: &mut [i16]| {
// SAFETY: `options.use_neon()` only returns true if neon is supported.
unsafe { neon::upsample_vertical_neon(a, b, c, d, e) }
};
}
#[cfg(feature = "portable_simd")]
return portable_simd::upsample_vertical_simd;
return scalar::upsample_vertical;
}
/// Upsample nothing
pub fn upsample_no_op(
_input: &[i16],
_in_ref: &[i16],
_in_near: &[i16],
_scratch_space: &mut [i16],
_output: &mut [i16],
) {
}
pub fn generic_sampler() -> UpSampler {
scalar::upsample_generic
}
#[cfg(test)]
mod tests {
use super::*;
#[cfg(feature = "portable_simd")]
mod portable_simd_impl {
use super::*;
#[test]
fn portable_simd_vertical() {
_test_vertical(portable_simd::upsample_vertical_simd)
}
#[test]
fn portable_simd_horizontal() {
_test_horizontal(portable_simd::upsample_horizontal_simd)
}
#[test]
fn portable_simd_hv() {
_test_hv(portable_simd::upsample_hv_simd)
}
}
#[cfg(any(target_arch = "x86", target_arch = "x86_64"))]
#[cfg(feature = "x86")]
#[cfg(target_feature = "avx2")]
mod avx2_impl {
use super::*;
#[test]
fn avx2_vertical() {
_test_vertical(|a: &[i16], b: &[i16], c: &[i16], d: &mut [i16], e: &mut [i16]| {
// SAFETY: Test guarded behind `target_feature`
unsafe { avx2::upsample_vertical_avx2(a, b, c, d, e) }
})
}
#[test]
fn avx2_horizontal() {
_test_horizontal(|a: &[i16], b: &[i16], c: &[i16], d: &mut [i16], e: &mut [i16]| {
// SAFETY: Test guarded behind `target_feature`
unsafe { avx2::upsample_horizontal_avx2(a, b, c, d, e) }
})
}
#[test]
fn avx2_hv() {
_test_hv(|a: &[i16], b: &[i16], c: &[i16], d: &mut [i16], e: &mut [i16]| {
// SAFETY: Test guarded behind `target_feature`
unsafe { avx2::upsample_hv_avx2(a, b, c, d, e) }
})
}
}
#[cfg(target_arch = "aarch64")]
#[cfg(feature = "neon")]
#[cfg(target_feature = "neon")]
mod neon_impl {
use super::*;
#[test]
fn neon_vertical() {
_test_vertical(|a: &[i16], b: &[i16], c: &[i16], d: &mut [i16], e: &mut [i16]| {
// SAFETY: Test guarded behind `target_feature`
unsafe { neon::upsample_vertical_neon(a, b, c, d, e) }
})
}
#[test]
fn neon_horizontal() {
_test_horizontal(|a: &[i16], b: &[i16], c: &[i16], d: &mut [i16], e: &mut [i16]| {
// SAFETY: Test guarded behind `target_feature`
unsafe { neon::upsample_horizontal_neon(a, b, c, d, e) }
})
}
#[test]
fn neon_hv() {
_test_hv(|a: &[i16], b: &[i16], c: &[i16], d: &mut [i16], e: &mut [i16]| {
// SAFETY: Test guarded behind `target_feature`
unsafe { neon::upsample_hv_neon(a, b, c, d, e) }
})
}
}
fn _test_vertical(upsampler: UpSampler) {
let width = 1024;
let input: Vec<i16> = (0..width).map(|x| ((x + 10) % 256) as i16).collect();
let in_near: Vec<i16> = (0..width).map(|x| ((x + 20) % 256) as i16).collect();
let in_far: Vec<i16> = (0..width).map(|x| ((x + 30) % 256) as i16).collect();
let mut scratch = vec![0i16; width];
let mut output_scalar = vec![0i16; width * 2];
let mut output_fast = vec![0i16; width * 2];
scalar::upsample_vertical(&input, &in_near, &in_far, &mut scratch, &mut output_scalar);
upsampler(&input, &in_near, &in_far, &mut scratch, &mut output_fast);
assert_eq!(output_scalar, output_fast);
}
fn _test_horizontal(upsampler: UpSampler) {
_test_horizontal_even_width(upsampler);
_test_horizontal_odd_width(upsampler);
}
fn _test_horizontal_even_width(upsampler: UpSampler) {
let width = 1024;
let input: Vec<i16> = (0..width).map(|x| ((x + 10) % 256) as i16).collect();
let mut scratch = vec![0i16; width];
let mut output_scalar = vec![0i16; width * 2];
let mut output_fast = vec![0i16; width * 2];
scalar::upsample_horizontal(&input, &[], &[], &mut scratch, &mut output_scalar);
upsampler(&input, &[], &[], &mut scratch, &mut output_fast);
assert_eq!(output_scalar, output_fast);
}
fn _test_horizontal_odd_width(upsampler: UpSampler) {
let width = 33;
let input: Vec<i16> = (0..width).map(|x| ((x + 10) % 256) as i16).collect();
let mut scratch = vec![0i16; width];
let mut output_scalar = vec![0i16; width * 2];
let mut output_fast = vec![0i16; width * 2];
scalar::upsample_horizontal(&input, &[], &[], &mut scratch, &mut output_scalar);
upsampler(&input, &[], &[], &mut scratch, &mut output_fast);
assert_eq!(output_scalar, output_fast);
}
fn _test_hv(upsampler: UpSampler) {
let width = 512;
let input: Vec<i16> = (0..width).map(|x| ((x + 10) % 256) as i16).collect();
let in_near: Vec<i16> = (0..width).map(|x| ((x + 20) % 256) as i16).collect();
let in_far: Vec<i16> = (0..width).map(|x| ((x + 30) % 256) as i16).collect();
// Output len is width * 4 for HV (vertical * 2, then horizontal * 2 for each row)
// scratch is width * 2
let mut scratch_scalar = vec![0i16; width * 2];
let mut scratch_fast = vec![0i16; width * 2];
let mut output_scalar = vec![0i16; width * 4];
let mut output_fast = vec![0i16; width * 4];
scalar::upsample_hv(
&input,
&in_near,
&in_far,
&mut scratch_scalar,
&mut output_scalar,
);
upsampler(
&input,
&in_near,
&in_far,
&mut scratch_fast,
&mut output_fast,
);
assert_eq!(output_scalar, output_fast);
}
}
+199
View File
@@ -0,0 +1,199 @@
/*
* Copyright (c) 2025.
*
* This software is free software;
*
* You can redistribute it or modify it under terms of the MIT, Apache License or Zlib license
*/
#[cfg(target_arch = "x86")]
use core::arch::x86::*;
#[cfg(target_arch = "x86_64")]
use core::arch::x86_64::*;
#[cfg(any(target_arch = "x86", target_arch = "x86_64"))]
#[target_feature(enable = "avx2")]
pub unsafe fn upsample_horizontal_avx2(
input: &[i16],
in_near: &[i16],
in_far: &[i16],
scratch: &mut [i16],
output: &mut [i16],
) {
assert_eq!(input.len() * 2, output.len());
assert!(input.len() > 2);
let len = input.len();
if len < 18 {
return super::scalar::upsample_horizontal(input, in_near, in_far, scratch, output);
}
// First two pixels
output[0] = input[0];
output[1] = (input[0] * 3 + input[1] + 2) >> 2;
let v_three = _mm256_set1_epi16(3);
let v_two = _mm256_set1_epi16(2);
let upsample16 = |input: &[i16; 18], output: &mut [i16; 32]| {
let in_ptr = input.as_ptr();
let out_ptr = output.as_mut_ptr();
// SAFETY: The input is 18 * 16 bit long, so the loads are safe.
let (v_prev, v_curr, v_next) = unsafe {
(
_mm256_loadu_si256(in_ptr.add(0) as *const __m256i),
_mm256_loadu_si256(in_ptr.add(1) as *const __m256i),
_mm256_loadu_si256(in_ptr.add(2) as *const __m256i),
)
};
let v_common = _mm256_add_epi16(_mm256_mullo_epi16(v_curr, v_three), v_two);
let v_even = _mm256_srai_epi16(_mm256_add_epi16(v_common, v_prev), 2);
let v_odd = _mm256_srai_epi16(_mm256_add_epi16(v_common, v_next), 2);
let v_res_1 = _mm256_unpacklo_epi16(v_even, v_odd);
let v_res_2 = _mm256_unpackhi_epi16(v_even, v_odd);
let v_final_1 = _mm256_permute2x128_si256(v_res_1, v_res_2, 0x20);
let v_final_2 = _mm256_permute2x128_si256(v_res_1, v_res_2, 0x31);
// SAFETY: The output is 32 * 16 bit long, so the stores are safe.
unsafe {
_mm256_storeu_si256(out_ptr as *mut __m256i, v_final_1);
_mm256_storeu_si256(out_ptr.add(16) as *mut __m256i, v_final_2);
}
};
for (input, output) in input
.windows(18)
.step_by(16)
.zip(output[2..].chunks_exact_mut(32))
{
upsample16(input.try_into().unwrap(), output.try_into().unwrap());
}
// Upsample the remainder. This may have some overlap, but that's fine.
if let Some(rest_input) = input.last_chunk::<18>() {
let end = output.len() - 2;
if let Some(rest_output) = output[..end].last_chunk_mut::<32>() {
upsample16(rest_input, rest_output);
}
}
// Last two pixels.
output[output.len() - 2] = (3 * input[len - 1] + input[len - 2] + 2) >> 2;
output[output.len() - 1] = input[len - 1];
}
#[cfg(any(target_arch = "x86", target_arch = "x86_64"))]
#[target_feature(enable = "avx2")]
pub unsafe fn upsample_vertical_avx2(
input: &[i16],
in_near: &[i16],
in_far: &[i16],
scratch: &mut [i16],
output: &mut [i16],
) {
assert_eq!(input.len() * 2, output.len());
assert_eq!(in_near.len(), input.len());
assert_eq!(in_far.len(), input.len());
let len = input.len();
if len < 16 {
return super::scalar::upsample_vertical(input, in_near, in_far, scratch, output);
}
let middle = output.len() / 2;
let (out_top, out_bottom) = output.split_at_mut(middle);
let v_three = _mm256_set1_epi16(3);
let v_two = _mm256_set1_epi16(2);
let upsample16 = |input: &[i16; 16],
in_near: &[i16; 16],
in_far: &[i16; 16],
out_top: &mut [i16; 16],
out_bottom: &mut [i16; 16]| {
// SAFETY: Inputs are all 16 * 16 bit long, so the loads are safe.
let (v_in, v_near, v_far) = unsafe {
(
_mm256_loadu_si256(input.as_ptr() as *const __m256i),
_mm256_loadu_si256(in_near.as_ptr() as *const __m256i),
_mm256_loadu_si256(in_far.as_ptr() as *const __m256i),
)
};
let v_common = _mm256_add_epi16(_mm256_mullo_epi16(v_in, v_three), v_two);
let v_out_top = _mm256_srai_epi16(_mm256_add_epi16(v_common, v_near), 2);
let v_out_bottom = _mm256_srai_epi16(_mm256_add_epi16(v_common, v_far), 2);
// SAFETY: Outputs are 16 * 16 bit long, so the stores are safe.
unsafe {
_mm256_storeu_si256(out_top.as_mut_ptr() as *mut __m256i, v_out_top);
_mm256_storeu_si256(out_bottom.as_mut_ptr() as *mut __m256i, v_out_bottom);
}
};
let chunks = input
.chunks_exact(16)
.zip(in_near.chunks_exact(16))
.zip(in_far.chunks_exact(16))
.zip(out_top.chunks_exact_mut(16))
.zip(out_bottom.chunks_exact_mut(16));
for ((((input, in_near), in_far), out_top), out_bottom) in chunks {
upsample16(
input.try_into().unwrap(),
in_near.try_into().unwrap(),
in_far.try_into().unwrap(),
out_top.try_into().unwrap(),
out_bottom.try_into().unwrap(),
);
}
// Upsample the remainder. This may have some overlap, but that's fine.
// Edition upgrade will fix this nested awfulness.
if let Some(rest) = input.last_chunk::<16>() {
if let Some(rest_near) = in_near.last_chunk::<16>() {
if let Some(rest_far) = in_far.last_chunk::<16>() {
if let Some(mut rest_top) = out_top.last_chunk_mut::<16>() {
if let Some(mut rest_bottom) = out_bottom.last_chunk_mut::<16>() {
upsample16(rest, rest_near, rest_far, &mut rest_top, &mut rest_bottom);
}
}
}
}
}
}
#[cfg(any(target_arch = "x86", target_arch = "x86_64"))]
#[target_feature(enable = "avx2")]
pub unsafe fn upsample_hv_avx2(
input: &[i16],
in_near: &[i16],
in_far: &[i16],
scratch_space: &mut [i16],
output: &mut [i16],
) {
assert_eq!(input.len() * 4, output.len());
assert!(input.len() * 2 <= scratch_space.len());
let scratch_space = &mut scratch_space[..input.len() * 2];
upsample_vertical_avx2(input, in_near, in_far, &mut [], scratch_space);
let scratch_half = scratch_space.len() / 2;
let output_half = output.len() / 2;
let (scratch_top, scratch_bottom) = scratch_space.split_at_mut(scratch_half);
let (out_top, out_bottom) = output.split_at_mut(output_half);
let mut t = [0];
upsample_horizontal_avx2(scratch_top, &[], &[], &mut t, out_top);
upsample_horizontal_avx2(scratch_bottom, &[], &[], &mut t, out_bottom);
}
+191
View File
@@ -0,0 +1,191 @@
/*
* Copyright (c) 2025.
*
* This software is free software;
*
* You can redistribute it or modify it under terms of the MIT, Apache License or Zlib license
*/
#[cfg(target_arch = "aarch64")]
use core::arch::aarch64::*;
#[cfg(target_arch = "aarch64")]
#[target_feature(enable = "neon")]
pub unsafe fn upsample_horizontal_neon(
input: &[i16], in_near: &[i16], in_far: &[i16], scratch: &mut [i16], output: &mut [i16]
) {
assert_eq!(input.len() * 2, output.len());
assert!(input.len() > 2);
let len = input.len();
if len < 10 {
return super::scalar::upsample_horizontal(input, in_near, in_far, scratch, output);
}
// First two pixels
output[0] = input[0];
output[1] = (input[0] * 3 + input[1] + 2) >> 2;
// SAFETY: NEON target feature is enabled on this function.
let v_three = unsafe { vdupq_n_s16(3) };
// SAFETY: NEON target feature is enabled on this function.
let v_two = unsafe { vdupq_n_s16(2) };
let upsample8 = |input: &[i16; 10], output: &mut [i16; 16]| {
let in_ptr = input.as_ptr();
let out_ptr = output.as_mut_ptr();
// SAFETY: The input is 10 * 16 bit long, so the loads are safe.
let (v_prev, v_curr, v_next) = unsafe {
(
vld1q_s16(in_ptr),
vld1q_s16(in_ptr.add(1)),
vld1q_s16(in_ptr.add(2))
)
};
// SAFETY: NEON target feature is enabled and vector lanes are valid.
let v_common = unsafe { vaddq_s16(vmulq_s16(v_curr, v_three), v_two) };
// SAFETY: NEON target feature is enabled and vector lanes are valid.
let v_even = unsafe { vshrq_n_s16::<2>(vaddq_s16(v_common, v_prev)) };
// SAFETY: NEON target feature is enabled and vector lanes are valid.
let v_odd = unsafe { vshrq_n_s16::<2>(vaddq_s16(v_common, v_next)) };
// SAFETY: NEON target feature is enabled and vector lanes are valid.
let v_res_1 = unsafe { vzip1q_s16(v_even, v_odd) };
// SAFETY: NEON target feature is enabled and vector lanes are valid.
let v_res_2 = unsafe { vzip2q_s16(v_even, v_odd) };
// SAFETY: The output is 16 * 16 bit long, so the stores are safe.
unsafe {
vst1q_s16(out_ptr, v_res_1);
vst1q_s16(out_ptr.add(8), v_res_2);
}
};
for (input, output) in input
.windows(10)
.step_by(8)
.zip(output[2..].chunks_exact_mut(16))
{
upsample8(input.try_into().unwrap(), output.try_into().unwrap());
}
// Upsample the remainder. This may have some overlap, but that's fine.
if let Some(rest_input) = input.last_chunk::<10>() {
let end = output.len() - 2;
if let Some(rest_output) = output[..end].last_chunk_mut::<16>() {
upsample8(rest_input, rest_output);
}
}
// Last two pixels.
output[output.len() - 2] = (3 * input[len - 1] + input[len - 2] + 2) >> 2;
output[output.len() - 1] = input[len - 1];
}
#[cfg(target_arch = "aarch64")]
#[target_feature(enable = "neon")]
pub unsafe fn upsample_vertical_neon(
input: &[i16], in_near: &[i16], in_far: &[i16], scratch: &mut [i16], output: &mut [i16]
) {
assert_eq!(input.len() * 2, output.len());
assert_eq!(in_near.len(), input.len());
assert_eq!(in_far.len(), input.len());
let len = input.len();
if len < 16 {
return super::scalar::upsample_vertical(input, in_near, in_far, scratch, output);
}
let middle = output.len() / 2;
let (out_top, out_bottom) = output.split_at_mut(middle);
// SAFETY: NEON target feature is enabled on this function.
let v_three = unsafe { vdupq_n_s16(3) };
// SAFETY: NEON target feature is enabled on this function.
let v_two = unsafe { vdupq_n_s16(2) };
let upsample8 = |input: &[i16; 8],
in_near: &[i16; 8],
in_far: &[i16; 8],
out_top: &mut [i16; 8],
out_bottom: &mut [i16; 8]| {
// SAFETY: Inputs are all 8 * 16 bit long, so the loads are safe.
let (v_in, v_near, v_far) = unsafe {
(
vld1q_s16(input.as_ptr()),
vld1q_s16(in_near.as_ptr()),
vld1q_s16(in_far.as_ptr())
)
};
// SAFETY: NEON target feature is enabled and vector lanes are valid.
let v_common = unsafe { vaddq_s16(vmulq_s16(v_in, v_three), v_two) };
// SAFETY: NEON target feature is enabled and vector lanes are valid.
let v_out_top = unsafe { vshrq_n_s16::<2>(vaddq_s16(v_common, v_near)) };
// SAFETY: NEON target feature is enabled and vector lanes are valid.
let v_out_bottom = unsafe { vshrq_n_s16::<2>(vaddq_s16(v_common, v_far)) };
// SAFETY: Outputs are 8 * 16 bit long, so the stores are safe.
unsafe {
vst1q_s16(out_top.as_mut_ptr(), v_out_top);
vst1q_s16(out_bottom.as_mut_ptr(), v_out_bottom);
}
};
let chunks = input
.chunks_exact(8)
.zip(in_near.chunks_exact(8))
.zip(in_far.chunks_exact(8))
.zip(out_top.chunks_exact_mut(8))
.zip(out_bottom.chunks_exact_mut(8));
for ((((input, in_near), in_far), out_top), out_bottom) in chunks {
upsample8(
input.try_into().unwrap(),
in_near.try_into().unwrap(),
in_far.try_into().unwrap(),
out_top.try_into().unwrap(),
out_bottom.try_into().unwrap()
);
}
// Upsample the remainder.
if let Some(rest) = input.last_chunk::<8>() {
if let Some(rest_near) = in_near.last_chunk::<8>() {
if let Some(rest_far) = in_far.last_chunk::<8>() {
if let Some(mut rest_top) = out_top.last_chunk_mut::<8>() {
if let Some(mut rest_bottom) = out_bottom.last_chunk_mut::<8>() {
upsample8(rest, rest_near, rest_far, &mut rest_top, &mut rest_bottom);
}
}
}
}
}
}
#[cfg(target_arch = "aarch64")]
#[target_feature(enable = "neon")]
pub unsafe fn upsample_hv_neon(
input: &[i16], in_near: &[i16], in_far: &[i16], scratch_space: &mut [i16], output: &mut [i16]
) {
assert_eq!(input.len() * 4, output.len());
assert!(input.len() * 2 <= scratch_space.len());
let scratch_space = &mut scratch_space[..input.len() * 2];
unsafe { upsample_vertical_neon(input, in_near, in_far, &mut [], scratch_space) };
let scratch_half = scratch_space.len() / 2;
let output_half = output.len() / 2;
let (scratch_top, scratch_bottom) = scratch_space.split_at_mut(scratch_half);
let (out_top, out_bottom) = output.split_at_mut(output_half);
let mut t = [0];
unsafe { upsample_horizontal_neon(scratch_top, &[], &[], &mut t, out_top) };
unsafe { upsample_horizontal_neon(scratch_bottom, &[], &[], &mut t, out_bottom) };
}
+171
View File
@@ -0,0 +1,171 @@
/*
* Copyright (c) 2023.
*
* This software is free software;
*
* You can redistribute it or modify it under terms of the MIT, Apache License or Zlib license
*/
use std::simd::prelude::*;
const LANES: usize = 16;
type V = Simd<i16, LANES>;
pub fn upsample_horizontal_simd(
input: &[i16],
in_near: &[i16],
in_far: &[i16],
scratch: &mut [i16],
output: &mut [i16],
) {
assert_eq!(input.len() * 2, output.len());
assert!(input.len() > 2);
let len = input.len();
if len < 18 {
return super::scalar::upsample_horizontal(input, in_near, in_far, scratch, output);
}
// First two pixels
output[0] = input[0];
output[1] = (input[0] * 3 + input[1] + 2) >> 2;
let v_three = V::splat(3);
let v_two = V::splat(2);
let upsample16 = |input: &[i16; 18], output: &mut [i16; 32]| {
let v_prev = V::from_slice(&input[0..LANES]);
let v_curr = V::from_slice(&input[1..LANES + 1]);
let v_next = V::from_slice(&input[2..LANES + 2]);
let v_common = v_curr * v_three + v_two;
let v_even = (v_common + v_prev) >> 2;
let v_odd = (v_common + v_next) >> 2;
let (v_res_1, v_res_2) = v_even.interleave(v_odd);
v_res_1.copy_to_slice(&mut output[0..LANES]);
v_res_2.copy_to_slice(&mut output[LANES..2 * LANES]);
};
for (input, output) in input
.windows(18)
.step_by(16)
.zip(output[2..].chunks_exact_mut(32))
{
upsample16(input.try_into().unwrap(), output.try_into().unwrap());
}
// Upsample the remainder. This may have some overlap, but that's fine.
if let Some(rest_input) = input.last_chunk::<18>() {
let end = output.len() - 2;
if let Some(rest_output) = output[..end].last_chunk_mut::<32>() {
upsample16(rest_input, rest_output);
}
}
// Last two pixels.
output[output.len() - 2] = (3 * input[len - 1] + input[len - 2] + 2) >> 2;
output[output.len() - 1] = input[len - 1];
}
pub fn upsample_vertical_simd(
input: &[i16],
in_near: &[i16],
in_far: &[i16],
_scratch_space: &mut [i16],
output: &mut [i16],
) {
assert_eq!(input.len() * 2, output.len());
assert_eq!(in_near.len(), input.len());
assert_eq!(in_far.len(), input.len());
let len = input.len();
if len < 16 {
return super::scalar::upsample_vertical(input, in_near, in_far, _scratch_space, output);
}
let middle = output.len() / 2;
let (out_top, out_bottom) = output.split_at_mut(middle);
let v_three = V::splat(3);
let v_two = V::splat(2);
let upsample16 = |input: &[i16; 16],
in_near: &[i16; 16],
in_far: &[i16; 16],
out_top: &mut [i16; 16],
out_bottom: &mut [i16; 16]| {
let v_in = V::from(*input);
let v_near = V::from(*in_near);
let v_far = V::from(*in_far);
let v_common = v_in * v_three + v_two;
let v_out_top = (v_common + v_near) >> 2;
let v_out_bottom = (v_common + v_far) >> 2;
v_out_top.copy_to_slice(out_top.as_mut_slice());
v_out_bottom.copy_to_slice(out_bottom.as_mut_slice());
};
let chunks = input
.chunks_exact(16)
.zip(in_near.chunks_exact(16))
.zip(in_far.chunks_exact(16))
.zip(out_top.chunks_exact_mut(16))
.zip(out_bottom.chunks_exact_mut(16));
for ((((input, in_near), in_far), out_top), out_bottom) in chunks {
upsample16(
input.try_into().unwrap(),
in_near.try_into().unwrap(),
in_far.try_into().unwrap(),
out_top.try_into().unwrap(),
out_bottom.try_into().unwrap(),
);
}
// Upsample the remainder. This may have some overlap, but that's fine.
// Edition upgrade will fix this nested awfulness.
if let Some(rest) = input.last_chunk::<16>() {
if let Some( rest_near) = in_near.last_chunk::<16>() {
if let Some( rest_far) = in_far.last_chunk::<16>() {
if let Some( rest_top) = out_top.last_chunk_mut::<16>() {
if let Some( rest_bottom) = out_bottom.last_chunk_mut::<16>() {
upsample16(rest, rest_near, rest_far, rest_top, rest_bottom);
}
}
}
}
}
}
pub fn upsample_hv_simd(
input: &[i16],
in_near: &[i16],
in_far: &[i16],
scratch_space: &mut [i16],
output: &mut [i16],
) {
assert_eq!(input.len() * 4, output.len());
assert!(input.len() * 2 <= scratch_space.len());
let scratch_space = &mut scratch_space[..input.len() * 2];
upsample_vertical_simd(input, in_near, in_far, &mut [], scratch_space);
let scratch_half = scratch_space.len() / 2;
let output_half = output.len() / 2;
let (scratch_top, scratch_bottom) = scratch_space.split_at_mut(scratch_half);
let (out_top, out_bottom) = output.split_at_mut(output_half);
let mut t = [0];
upsample_horizontal_simd(scratch_top, &[], &[], &mut t, out_top);
upsample_horizontal_simd(scratch_bottom, &[], &[], &mut t, out_bottom);
}
+129
View File
@@ -0,0 +1,129 @@
/*
* Copyright (c) 2023.
*
* This software is free software;
*
* You can redistribute it or modify it under terms of the MIT, Apache License or Zlib license
*/
pub fn upsample_horizontal(
input: &[i16], _ref: &[i16], _in_near: &[i16], _scratch: &mut [i16], output: &mut [i16]
) {
assert_eq!(
input.len() * 2,
output.len(),
"Input length is not half the size of the output length"
);
assert!(
output.len() > 4 && input.len() > 2,
"Too Short of a vector, cannot upsample"
);
output[0] = input[0];
output[1] = (input[0] * 3 + input[1] + 2) >> 2;
// This code is written for speed and not readability
//
// The readable code is
//
// for i in 1..input.len() - 1{
// let sample = 3 * input[i] + 2;
// out[i * 2] = (sample + input[i - 1]) >> 2;
// out[i * 2 + 1] = (sample + input[i + 1]) >> 2;
// }
//
// The output of a pixel is determined by it's surrounding neighbours but we attach more weight to it's nearest
// neighbour (input[i]) than to the next nearest neighbour.
for (output_window, input_window) in output[2..].chunks_exact_mut(2).zip(input.windows(3)) {
let sample = 3 * input_window[1] + 2;
output_window[0] = (sample + input_window[0]) >> 2;
output_window[1] = (sample + input_window[2]) >> 2;
}
// Get lengths
let out_len = output.len() - 2;
let input_len = input.len() - 2;
// slice the output vector
let f_out = &mut output[out_len..];
let i_last = &input[input_len..];
// write out manually..
f_out[0] = (3 * i_last[1] + i_last[0] + 2) >> 2;
f_out[1] = i_last[1];
}
pub fn upsample_vertical(
input: &[i16], in_near: &[i16], in_far: &[i16], _scratch_space: &mut [i16], output: &mut [i16]
) {
assert_eq!(input.len() * 2, output.len());
assert_eq!(in_near.len(), input.len());
assert_eq!(in_far.len(), input.len());
let middle = output.len() / 2;
let (out_top, out_bottom) = output.split_at_mut(middle);
// for the first row, closest row is in_near
for ((near, far), x) in input.iter().zip(in_near.iter()).zip(out_top) {
*x = (((3 * near) + 2) + far) >> 2;
}
// for the second row, the closest row to input is in_far
for ((near, far), x) in input.iter().zip(in_far.iter()).zip(out_bottom) {
*x = (((3 * near) + 2) + far) >> 2;
}
}
pub fn upsample_hv(
input: &[i16], in_near: &[i16], in_far: &[i16], scratch_space: &mut [i16], output: &mut [i16]
) {
assert_eq!(input.len() * 4, output.len());
assert!(input.len() * 2 <= scratch_space.len());
let scratch_space = &mut scratch_space[..input.len() * 2];
let mut t = [0];
upsample_vertical(input, in_near, in_far, &mut t, scratch_space);
// horizontal upsampling must be done separate for every line
// Otherwise it introduces artifacts that may cause the edge colors
// to appear on the other line.
// Since this is called for two scanlines/widths currently
// splitting the inputs and outputs into half ensures we only handle
// one scanline per iteration
let scratch_half = scratch_space.len() / 2;
let output_half = output.len() / 2;
upsample_horizontal(
&scratch_space[..scratch_half],
&[],
&[],
&mut t,
&mut output[..output_half]
);
upsample_horizontal(
&scratch_space[scratch_half..],
&[],
&[],
&mut t,
&mut output[output_half..]
);
}
pub fn upsample_generic(
input: &[i16], _in_near: &[i16], _in_far: &[i16], _scratch_space: &mut [i16],
output: &mut [i16]
) {
// use nearest sample
let difference = output.len() / input.len();
if difference > 0 {
// nearest neighbour
for (input, chunk_output) in input.iter().zip(output.chunks_exact_mut(difference)) {
chunk_output.iter_mut().for_each(|x| *x = *input);
}
}
}

Some files were not shown because too many files have changed in this diff Show More