From 2b33d9f6f043c08d79d03ceabdd2a03af6fdfb28 Mon Sep 17 00:00:00 2001 From: RoyLin Date: Sun, 19 Jul 2026 05:29:13 +0800 Subject: [PATCH 1/9] feat(ocr): add built-in first-use capability --- .github/workflows/release.yml | 48 +- Cargo.lock | 20 + Cargo.toml | 7 +- README.md | 117 +++- .../skills/a3s-use-browser/SKILL.md | 5 +- crates/browser-driver/src/mcp.rs | 4 + crates/extension/src/lib.rs | 5 +- crates/ocr/Cargo.toml | 34 + crates/ocr/README.md | 34 + crates/ocr/skills/a3s-use-ocr/SKILL.md | 48 ++ crates/ocr/src/cli.rs | 205 ++++++ crates/ocr/src/client.rs | 649 ++++++++++++++++++ crates/ocr/src/lib.rs | 18 + crates/ocr/src/main.rs | 39 ++ crates/ocr/src/mcp.rs | 170 +++++ crates/ocr/src/models.rs | 105 +++ crates/ocr/src/provider.rs | 393 +++++++++++ crates/office/skills/a3s-use-office/SKILL.md | 25 +- .../skills/a3s-use-office/references/mcp.md | 15 +- docs/architecture.md | 43 +- src/capability_registry.rs | 154 ++++- src/cli.rs | 88 ++- src/cli_tests.rs | 50 +- src/lib.rs | 6 + src/mcp.rs | 131 +++- src/mcp/office.rs | 96 ++- src/mcp/office/tests.rs | 26 +- src/ocr_builtin.rs | 95 +++ tests/cli.rs | 77 ++- 29 files changed, 2608 insertions(+), 99 deletions(-) create mode 100644 crates/ocr/Cargo.toml create mode 100644 crates/ocr/README.md create mode 100644 crates/ocr/skills/a3s-use-ocr/SKILL.md create mode 100644 crates/ocr/src/cli.rs create mode 100644 crates/ocr/src/client.rs create mode 100644 crates/ocr/src/lib.rs create mode 100644 crates/ocr/src/main.rs create mode 100644 crates/ocr/src/mcp.rs create mode 100644 crates/ocr/src/models.rs create mode 100644 crates/ocr/src/provider.rs create mode 100644 src/ocr_builtin.rs diff --git a/.github/workflows/release.yml b/.github/workflows/release.yml index dfee133e..0116089b 100644 --- a/.github/workflows/release.yml +++ b/.github/workflows/release.yml @@ -42,7 +42,7 @@ jobs: test "${{ inputs.release_tag }}" = "v${version}" fi - run: cargo fmt --all -- --check - - run: cargo test --workspace --all-features --locked + - run: cargo test --workspace --all-features --locked -- --test-threads=1 - run: cargo clippy --workspace --all-targets --all-features --locked -- -D warnings binaries: @@ -84,12 +84,13 @@ jobs: version="${{ needs.validate.outputs.version }}" archive="a3s-use-${version}-${{ matrix.name }}.tar.gz" stage="${RUNNER_TEMP}/a3s-use-${version}-${{ matrix.name }}" - install -d "${stage}/skills" "${stage}/skill-data" "${stage}/office-skills" "${stage}/dashboard" + install -d "${stage}/skills" "${stage}/skill-data" "${stage}/office-skills" "${stage}/ocr-skills" "${stage}/dashboard" install -m 0755 "target/${{ matrix.target }}/release/a3s-use" "${stage}/a3s-use" install -m 0755 "target/${{ matrix.target }}/release/a3s-use-browser-driver" "${stage}/a3s-use-browser-driver" cp -R crates/browser-driver/skills/. "${stage}/skills/" cp -R crates/browser-driver/skill-data/. "${stage}/skill-data/" cp -R crates/office/skills/. "${stage}/office-skills/" + cp -R crates/ocr/skills/. "${stage}/ocr-skills/" cp -R crates/browser-driver/dashboard/out/. "${stage}/dashboard/" install -m 0644 LICENSE README.md THIRD_PARTY_NOTICES.md "${stage}/" install -m 0644 crates/browser-driver/LICENSE-APACHE-2.0 "${stage}/LICENSE-APACHE-2.0" @@ -110,6 +111,7 @@ jobs: Copy-Item -Recurse "crates/browser-driver/skills" "$stage/skills" Copy-Item -Recurse "crates/browser-driver/skill-data" "$stage/skill-data" Copy-Item -Recurse "crates/office/skills" "$stage/office-skills" + Copy-Item -Recurse "crates/ocr/skills" "$stage/ocr-skills" Copy-Item -Recurse "crates/browser-driver/dashboard/out" "$stage/dashboard" Copy-Item LICENSE,README.md,THIRD_PARTY_NOTICES.md $stage Copy-Item crates/browser-driver/LICENSE-APACHE-2.0 "$stage/LICENSE-APACHE-2.0" @@ -128,6 +130,7 @@ jobs: test -x "${install_root}/a3s-use-browser-driver" test -f "${install_root}/skill-data/core/SKILL.md" test -f "${install_root}/office-skills/a3s-use-office/SKILL.md" + test -f "${install_root}/ocr-skills/a3s-use-ocr/SKILL.md" test -f "${install_root}/dashboard/index.html" test -f "${install_root}/LICENSE-APACHE-2.0" test -f "${install_root}/UPSTREAM.md" @@ -143,6 +146,29 @@ jobs: "agentcore", "core", "dogfood", "electron", "slack", "vercel-sandbox" } PY + "${install_root}/a3s-use" ocr doctor --json > "${RUNNER_TEMP}/ocr-doctor.json" + python3 - "${RUNNER_TEMP}/ocr-doctor.json" <<'PY' + import json, pathlib, sys + value = json.loads(pathlib.Path(sys.argv[1]).read_text()) + assert value["ok"] is True + assert value["data"]["readiness"] in {"ready", "missing", "broken", "unknown"} + PY + "${install_root}/a3s-use" capability snapshot --json > "${RUNNER_TEMP}/capabilities.json" + python3 - "${RUNNER_TEMP}/capabilities.json" "${install_root}" <<'PY' + import json, pathlib, sys + value = json.loads(pathlib.Path(sys.argv[1]).read_text()) + root = pathlib.Path(sys.argv[2]).resolve() + ocr = next( + capability + for capability in value["data"]["registry"]["capabilities"] + if capability["id"] == "use/ocr" + ) + assert ocr["mcp"]["target"] == "ocr-native" + assert len(ocr["skills"]) == 1 + assert pathlib.Path(ocr["skills"][0]["path"]).resolve() == ( + root / "ocr-skills" / "a3s-use-ocr" / "SKILL.md" + ) + PY A3S_OFFICECLI_EXECUTABLE="${install_root}/must-not-be-invoked" \ "${install_root}/a3s-use" office skills list --json > "${RUNNER_TEMP}/office-skills.json" python3 - "${RUNNER_TEMP}/office-skills.json" <<'PY' @@ -179,6 +205,7 @@ jobs: "$root/a3s-use-browser-driver.exe", "$root/skill-data/core/SKILL.md", "$root/office-skills/a3s-use-office/SKILL.md", + "$root/ocr-skills/a3s-use-ocr/SKILL.md", "$root/dashboard/index.html", "$root/LICENSE-APACHE-2.0", "$root/UPSTREAM.md", @@ -194,6 +221,19 @@ jobs: if (-not $officeSkills.ok -or $officeSkills.data.Count -ne 1 -or $officeSkills.data[0].name -ne "a3s-use-office") { throw "Packaged Office Skill smoke failed" } + $ocr = (& "$root/a3s-use.exe" ocr doctor --json | ConvertFrom-Json) + if (-not $ocr.ok -or -not $ocr.data.readiness) { throw "Built-in OCR doctor smoke failed" } + $capabilities = (& "$root/a3s-use.exe" capability snapshot --json | ConvertFrom-Json) + $ocrCapability = $capabilities.data.registry.capabilities | + Where-Object id -eq "use/ocr" + if ($ocrCapability.mcp.target -ne "ocr-native" -or $ocrCapability.skills.Count -ne 1) { + throw "Built-in OCR capability projection smoke failed" + } + $expectedOcrSkill = (Resolve-Path "$root/ocr-skills/a3s-use-ocr/SKILL.md").Path + $actualOcrSkill = (Resolve-Path $ocrCapability.skills[0].path).Path + if (-not [StringComparer]::OrdinalIgnoreCase.Equals($actualOcrSkill, $expectedOcrSkill)) { + throw "Built-in OCR Skill was not projected from the installed release" + } $requests = @( '{"jsonrpc":"2.0","id":1,"method":"initialize","params":{"protocolVersion":"2025-06-18","capabilities":{},"clientInfo":{"name":"release-smoke","version":"1"}}}', '{"jsonrpc":"2.0","method":"notifications/initialized","params":{}}', @@ -223,7 +263,7 @@ jobs: ref: ${{ github.event_name == 'workflow_dispatch' && inputs.release_tag || github.ref }} - uses: dtolnay/rust-toolchain@stable - uses: Swatinem/rust-cache@v2 - - name: Publish Core then Browser + - name: Publish Core, OCR, then Browser env: CARGO_REGISTRY_TOKEN: ${{ secrets.CARGO_TOKEN }} VERSION: ${{ needs.validate.outputs.version }} @@ -262,6 +302,8 @@ jobs: publish_once a3s-use-core wait_until_visible a3s-use-core + publish_once a3s-use-ocr + wait_until_visible a3s-use-ocr publish_once a3s-use-browser release: diff --git a/Cargo.lock b/Cargo.lock index 14cb1eb3..80e0f846 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -14,6 +14,7 @@ dependencies = [ "a3s-use-browser", "a3s-use-core", "a3s-use-extension", + "a3s-use-ocr", "a3s-use-office", "anyhow", "async-trait", @@ -115,6 +116,25 @@ dependencies = [ "tokio", ] +[[package]] +name = "a3s-use-ocr" +version = "0.1.1" +dependencies = [ + "a3s-use-core", + "axum", + "base64", + "clap", + "reqwest", + "rmcp", + "schemars", + "serde", + "serde_json", + "sha2 0.10.9", + "tempfile", + "tokio", + "url", +] + [[package]] name = "a3s-use-office" version = "0.1.1" diff --git a/Cargo.toml b/Cargo.toml index a136a13b..8c17bb9f 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -5,6 +5,7 @@ members = [ "crates/browser-driver", "crates/office", "crates/extension", + "crates/ocr", ] resolver = "2" @@ -48,7 +49,7 @@ license.workspace = true repository.workspace = true authors.workspace = true rust-version = "1.85" -description = "Typed Browser, Office, and external application capabilities for A3S" +description = "Typed Browser, Office, OCR, and external application capabilities for A3S" [lib] name = "a3s_use" @@ -59,7 +60,7 @@ name = "a3s-use" path = "src/main.rs" [features] -default = ["browser", "office", "extensions", "mcp"] +default = ["browser", "office", "ocr", "extensions", "mcp"] browser = ["dep:a3s-use-browser"] office = [ "dep:a3s-use-office", @@ -67,6 +68,7 @@ office = [ "dep:futures-util", "dep:getrandom", ] +ocr = ["dep:a3s-use-ocr"] extensions = ["dep:a3s-use-extension"] mcp = [ "dep:axum", @@ -84,6 +86,7 @@ lightpanda = ["browser", "a3s-use-browser/lightpanda"] a3s-use-core = { version = "0.1.1", path = "crates/core" } a3s-use-browser = { version = "0.1.1", path = "crates/browser", optional = true } a3s-use-office = { version = "0.1.1", path = "crates/office", optional = true } +a3s-use-ocr = { version = "0.1.1", path = "crates/ocr", optional = true } a3s-use-extension = { version = "0.1.1", path = "crates/extension", optional = true } anyhow.workspace = true axum = { workspace = true, optional = true } diff --git a/README.md b/README.md index 9b8376a1..5b4916fa 100644 --- a/README.md +++ b/README.md @@ -5,7 +5,7 @@

- Use browsers, Office documents, and independently shipped application domains through native CLI, standard MCP, and Skills + Use browsers, Office documents, OCR, and independently shipped application domains through native CLI, standard MCP, and Skills

@@ -14,6 +14,7 @@ Quick Start • Browser • Office • + OCR • Extensions • Architecture • Development @@ -23,10 +24,10 @@ ## Overview -**A3S Use** is the application-capability layer for A3S. Browser and Office are -first-party domains in the default distribution. Independently distributed -packages can add more domains without rebuilding Use by declaring native CLI, -standard MCP, and/or `SKILL.md` surfaces in an A3S ACL manifest. +**A3S Use** is the application-capability layer for A3S. Browser, native Office, +and OCR are first-party domains in the default distribution. Independently +distributed packages can add more domains without rebuilding Use by declaring +native CLI, standard MCP, and/or `SKILL.md` surfaces in an A3S ACL manifest. The primary user entry point is `a3s use`; `a3s-use` is the standalone binary used by the umbrella CLI and remains available for direct use, automation, and @@ -95,7 +96,14 @@ a3s use mcp serve browser a3s use mcp serve office-native # Keep using the pinned OfficeCLI compatibility MCP server where needed. +a3s use mcp serve office-compat +# Legacy alias: a3s use mcp serve office + +# Built-in OCR; provider readiness remains explicit. +a3s use ocr doctor --json +a3s use ocr extract ./scan.png --language eng --json +a3s use mcp serve ocr ``` Every domain argument accepted by `a3s use ...` can also be passed directly to @@ -103,8 +111,8 @@ Every domain argument accepted by `a3s use ...` can also be passed directly to ## Features -- **Built-In Browser and Office**: Keep stable first-party command routes while - reporting provider readiness separately +- **Built-In Browser, Office, and OCR**: Keep stable first-party command routes + while reporting provider readiness separately - **Typed Rust Contracts**: Embed Browser rendering and Office operations without starting a CLI process or an MCP server - **Agent Browser Compatibility**: Provide the locked 82-command vocabulary, @@ -129,6 +137,9 @@ Every domain argument accepted by `a3s use ...` can also be passed directly to safe Word, Spreadsheet, Presentation, native MCP, and compatibility workflows - **External Domains**: Install process-isolated packages that expose any useful combination of CLI, MCP, and Skill surfaces +- **First-Party OCR Domain**: Extract text and bounded layout evidence with + a local Tesseract provider or an explicitly configured vision endpoint, + without silently installing a provider or hiding remote image transfer - **Hot-Plug Discovery**: Publish immutable generation/revision snapshots so a resident host can add, replace, or remove live capabilities without restarting - **Content-Bound Skills**: Project an absolute package path and lowercase @@ -149,6 +160,7 @@ Every domain argument accepted by `a3s use ...` can also be passed directly to | Browser | Built in | Full Browser vocabulary | A3S Use standard MCP server | Six packaged Browser Skills | A3S Use | | Office | Built in | Stable Office vocabulary | Typed native preview plus OfficeCLI compatibility server | Packaged `a3s-use-office` Skill | A3S Use native engine; OfficeCLI compatibility in 0.1.x | | Box | Reserved built-in route | Native A3S Box vocabulary | — | — | Umbrella A3S CLI | +| OCR | Built in | Doctor and typed image extraction | `ocr_doctor` and `ocr_extract` | One provider-safe OCR Skill | A3S Use process and explicitly configured provider | | External domain | Installed extension | Optional native executable | Optional standard MCP server | Optional `SKILL.md` | Extension package plus A3S Use lifecycle | The Box route is component-backed. The umbrella CLI resolves its authoritative @@ -157,12 +169,13 @@ does not copy Box, discover a replacement on `PATH`, or write a second receipt. ### Cargo feature matrix -Default features are `browser`, `office`, `extensions`, and `mcp`. +Default features are `browser`, `office`, `ocr`, `extensions`, and `mcp`. | Feature | Included capability | | --- | --- | | `browser` | Typed Browser library, stateless rendering, and full Browser driver delegation | | `office` | Typed Office contracts, native OOXML read engine, and temporary OfficeCLI compatibility | +| `ocr` | Built-in typed OCR CLI/MCP with local Tesseract and explicit vision providers | | `extensions` | ACL manifests, package receipts, hot-plug registry, and external CLI/MCP/Skill routes | | `mcp` | Standard MCP servers plus the managed Browser Streamable HTTP lifecycle | | `lightpanda` | Explicit opt-in Lightpanda provider support in addition to Chrome | @@ -179,6 +192,7 @@ A compiled command surface is not proof that its provider is installed. Use | `a3s-use-browser-driver` | Complete interactive Browser CLI, MCP tools, Skills, Dashboard, and compatibility runtime | | `a3s-use-office` | Native OOXML foundation, typed Office operations, and compatibility lifecycle | | `a3s-use-extension` | A3S ACL manifest model, package registry, leases, and native surface descriptors | +| `a3s-use-ocr` | Typed local/vision OCR providers, CLI, MCP tools, and release-packaged Skill assets | | `a3s-use` | Facade library, standalone CLI host, capability projection, and MCP entry points | ## Quick Start @@ -198,9 +212,9 @@ a3s use doctor --json Prebuilt archives are also published on [GitHub Releases](https://github.com/A3S-Lab/Use/releases). A complete archive contains `a3s-use`, its sibling `a3s-use-browser-driver`, Browser Skills, the -first-party Office Skill, the Dashboard, and license/provenance notices. Keep -those packaged assets together; installing only the facade binary does not -provide the complete Browser and Office Skill surfaces. +first-party Office and OCR Skills, the Dashboard, and license/provenance +notices. Keep those packaged assets together; installing only the facade binary +does not provide the complete Browser, Office, and OCR Skill surfaces. Build all binaries from source with: @@ -515,7 +529,10 @@ agents without starting OfficeCLI. Discover its metadata with `office skills get a3s-use-office`, append its four format/MCP references with `--full`, or locate the installed directory with `office skills path`. The capability snapshot binds the Skill path and lowercase SHA-256 so a resident -host can verify the bytes before loading them. +host can verify the bytes before loading them. Resident Code hosts receive the +native engine as canonical route `use/office` targeting `office-native`; a ready +OfficeCLI installation is projected separately as `use/office-compat` targeting +`office-compat`. Other `0.1.x` commands and the default `mcp serve office` target still use a compatibility backend pinned to OfficeCLI `1.0.136`. This is a migration boundary, not a native-promotion claim. The default routes will be promoted @@ -1586,6 +1603,34 @@ compatibility response can return See [Native Office Engine](docs/native-office.md) for the complete requirements, compatibility scope, safety invariants, delivery gates, and migration plan. +## OCR + +`a3s-use-ocr` implements the reserved built-in `ocr` route. The default Use +release packages its `a3s-use-ocr` Skill and exposes `ocr_doctor` plus +`ocr_extract` over standard MCP, so a resident A3S Code session receives +`mcp__use_ocr__*` without installing a separate extension. + +OCR never installs a provider silently. `auto` prefers an explicitly configured +or discoverable Tesseract executable. Vision OCR is enabled only when its model +and endpoint configuration are present; non-loopback endpoints require HTTPS +and an API key, and the diagnostic discloses that the complete source image +leaves the device. Supported inputs are bounded local PNG, JPEG, WebP, GIF, +BMP, and TIFF files. The result binds the canonical source path, media type, +byte length, and SHA-256 alongside text and any available +confidence/bounding-box evidence. + +```bash +a3s use ocr doctor --json +a3s use ocr extract ./scan.png --language eng --json +a3s use mcp serve ocr +``` + +A3S Code may first-use install the verified parent Use release. OCR provider +selection remains explicit, and remote vision extraction still escalates to +the parent TUI before source bytes leave the device. + +See the [OCR crate](crates/ocr/README.md) for configuration and provider +boundaries. ## External Extensions External Use domains stay behind process boundaries. A package contains an @@ -1636,15 +1681,17 @@ roadmap work; Use does not silently install arbitrary Homebrew, npm, Cargo, system, or `PATH` packages. Built-in and management routes are reserved. Extensions cannot shadow -`browser`, `office`, `box`, `component`, `capability`, or other host commands. +`browser`, `office`, `ocr`, `box`, `component`, `capability`, or other host +commands. ## Live Host Integration Resident hosts consume `capability snapshot` and `capability watch`. The -projection presents Browser, Office, Box, and enabled extensions through one -read-only schema while preserving each binding's `built-in` or `extension` -origin. The extension generation advances on receipt mutations; a content -revision also changes when built-in readiness or packaged Skill content changes. +projection presents Browser, native Office, OCR, Box, and enabled extensions +through one read-only schema while preserving each binding's `built-in` or +`extension` origin. The extension generation advances on receipt mutations; a +content revision also changes when built-in readiness or packaged Skill content +changes. ```bash a3s-use capability snapshot --json @@ -1662,10 +1709,18 @@ tools. Projected Skills provide guidance only and cannot expand permissions or authorize installation. Code verifies their projected SHA-256 before loading the exact bytes. +The built-in Office projection is intentionally host-oriented: `use/office` +always exposes the in-process native MCP target when MCP support is compiled, +without consulting OfficeCLI. A discovered OfficeCLI provider is a separate +optional `use/office-compat` route, so native readiness and compatibility +installation cannot mask or replace each other. + A capability becomes callable only after its MCP connection is ready. A removed or replaced route leaves the worker catalog before its old connection -drains. Starting Code never installs Use: component installation remains an -explicit umbrella CLI action. +drains. Code TUI resolves the catalogued Use component on first launch and may +install its verified release before terminal takeover. Offline mode and +`A3S_NO_AUTO_INSTALL=1` remain strict no-mutation boundaries; setup failure is +non-fatal and stays visible through `/use`. ## Protocol and Lifecycle Boundaries @@ -1702,13 +1757,13 @@ crash, and in-flight calls retain the exact package generation they accepted. a3s use │ a3s-use host - ┌─────────────┼──────────────┐ - │ │ │ - Browser Office extension registry - typed + driver native OOXML CLI / MCP / Skill - + 0.1 compat - │ │ │ - └──────── capability snapshot/watch ───────► A3S Code + ┌──────────┬──────────┬──────────┬──────────────┐ + │ │ │ │ │ + Browser Office OCR extension registry + typed + driver OOXML local/vision CLI / MCP / Skill + + 0.1 compat + │ │ │ │ + └──────── capability snapshot/watch ───────────► A3S Code a3s-search ── Arc ──► a3s-use-browser @@ -1717,10 +1772,12 @@ crash, and in-flight calls retain the exact package generation they accepted. The dependency arrows are intentional. Search links only the Browser contract, so rendering does not require `a3s-use`, MCP, or a resident process. Office is -an in-process typed engine with a temporary 0.1.x compatibility process; -external domains retain their process boundaries. A3S Code consumes the -read-only projection and connects standard MCP/Skill surfaces; it does not gain -component installation authority. +an in-process typed engine with a temporary 0.1.x compatibility process. OCR +uses an explicitly present local Tesseract executable or an explicitly +configured vision provider; it never installs either silently. External +domains retain their process boundaries. A3S Code consumes the read-only +projection and connects standard MCP/Skill surfaces; bounded provider +installation requests still require the parent TUI's authority. Source is split between the facade under `src/` and focused workspace crates under `crates/`. See [Architecture](docs/architecture.md) for package leases, diff --git a/crates/browser-driver/skills/a3s-use-browser/SKILL.md b/crates/browser-driver/skills/a3s-use-browser/SKILL.md index 072d4ef5..ee03a73a 100644 --- a/crates/browser-driver/skills/a3s-use-browser/SKILL.md +++ b/crates/browser-driver/skills/a3s-use-browser/SKILL.md @@ -10,7 +10,10 @@ Use the host surface that is already available: - In an A3S Code `use` worker, call the available `mcp__use_browser__*` tools directly. The host owns installation and MCP - lifecycle; do not run component installation or shell commands there. + lifecycle; do not run component installation or shell commands there. Call + `mcp__use_browser__agent_browser_doctor` first. If its managed browser is + missing, request `mcp__use_browser__agent_browser_install`; the parent TUI + must obtain HITL approval before that mutation can run. - In a CLI-only agent host, use the `a3s use browser ...` commands below. Install the built-in capability and its managed runtime when needed: diff --git a/crates/browser-driver/src/mcp.rs b/crates/browser-driver/src/mcp.rs index 99fb2a30..c1585c07 100644 --- a/crates/browser-driver/src/mcp.rs +++ b/crates/browser-driver/src/mcp.rs @@ -352,6 +352,8 @@ const CORE_PROFILE_TOOLS: &[&str] = &[ TOOL_TAB_CLOSE, TOOL_EVAL, TOOL_CLOSE, + TOOL_DOCTOR, + TOOL_INSTALL, ]; const NETWORK_PROFILE_TOOLS: &[&str] = &[ @@ -3744,6 +3746,8 @@ mod tests { assert!(names.contains(&TOOL_SNAPSHOT)); assert!(names.contains(&TOOL_CLICK)); assert!(names.contains(&TOOL_SCREENSHOT)); + assert!(names.contains(&TOOL_DOCTOR)); + assert!(names.contains(&TOOL_INSTALL)); assert!(names.contains(&TOOL_GET_CDP_URL)); assert!(names.contains(&TOOL_NETWORK_HAR_START)); assert!(names.contains(&TOOL_REACT_SUSPENSE)); diff --git a/crates/extension/src/lib.rs b/crates/extension/src/lib.rs index ffc918a5..cc76e2dc 100644 --- a/crates/extension/src/lib.rs +++ b/crates/extension/src/lib.rs @@ -23,6 +23,9 @@ const RESERVED_ROUTES: &[&str] = &[ "box", "capability", "office", + "office-compat", + "office-native", + "ocr", "capabilities", "component", "extension", @@ -437,7 +440,7 @@ extension "acme/slack" { #[test] fn rejects_reserved_routes() { - for route in ["browser", "box"] { + for route in ["browser", "box", "ocr"] { let manifest = MANIFEST.replace( "route = \"slack\"", &format!("route = \"{route}\""), diff --git a/crates/ocr/Cargo.toml b/crates/ocr/Cargo.toml new file mode 100644 index 00000000..4a5f58aa --- /dev/null +++ b/crates/ocr/Cargo.toml @@ -0,0 +1,34 @@ +[package] +name = "a3s-use-ocr" +version.workspace = true +edition.workspace = true +license.workspace = true +repository.workspace = true +authors.workspace = true +rust-version.workspace = true +description = "Typed built-in optical character recognition for A3S Use" + +[lib] +name = "a3s_use_ocr" +path = "src/lib.rs" + +[[bin]] +name = "a3s-use-ocr" +path = "src/main.rs" + +[dependencies] +a3s-use-core = { version = "0.1.1", path = "../core" } +base64.workspace = true +clap.workspace = true +reqwest = { workspace = true, features = ["json"] } +rmcp.workspace = true +schemars.workspace = true +serde.workspace = true +serde_json.workspace = true +sha2.workspace = true +tokio.workspace = true +url.workspace = true + +[dev-dependencies] +axum.workspace = true +tempfile.workspace = true diff --git a/crates/ocr/README.md b/crates/ocr/README.md new file mode 100644 index 00000000..d61b4096 --- /dev/null +++ b/crates/ocr/README.md @@ -0,0 +1,34 @@ +# A3S Use OCR + +`a3s-use-ocr` implements the first-party built-in OCR domain for A3S Use. A3S +Code receives it as `mcp__use_ocr__*` through the release-matched Use registry, +without a separate extension install. It exposes the same typed extraction +through a native CLI and standard stdio MCP, and does not silently install an +OCR provider. + +Provider selection is explicit: + +- `A3S_OCR_PROVIDER=auto|tesseract|vision` +- `A3S_OCR_TESSERACT_EXECUTABLE=/absolute/path/to/tesseract` +- `A3S_OCR_VISION_MODEL=` +- `A3S_OCR_VISION_BASE_URL=https://provider.example/v1/` +- `A3S_OCR_VISION_API_KEY=` +- `A3S_OCR_TIMEOUT_MS=60000` + +`auto` prefers a configured or discoverable local Tesseract executable. It uses +the vision provider only when the vision environment is configured. Remote +vision endpoints require HTTPS and an API key; loopback HTTP is allowed for a +local provider. + +Build and exercise the domain through the Use facade: + +```bash +a3s use ocr doctor --json +a3s use ocr extract ./scan.png --language eng --json +a3s use mcp serve ocr +``` + +The A3S Use release packages the OCR Skill beside the facade binary. A3S Code +can first-use install that verified release and hot-plug the built-in route. +Provider setup remains explicit: local Tesseract never sends source bytes +off-device, while a configured remote vision provider requires parent HITL. diff --git a/crates/ocr/skills/a3s-use-ocr/SKILL.md b/crates/ocr/skills/a3s-use-ocr/SKILL.md new file mode 100644 index 00000000..9ca0950b --- /dev/null +++ b/crates/ocr/skills/a3s-use-ocr/SKILL.md @@ -0,0 +1,48 @@ +--- +name: a3s-use-ocr +description: Extract text and layout evidence from local image files through the built-in A3S Use OCR domain. Use when an agent needs optical character recognition for a PNG, JPEG, WebP, GIF, BMP, or TIFF image and must preserve the source digest, provider disclosure, confidence, and bounding-box evidence. +--- + +# A3S Use OCR + +Use the host-provided A3S Use surface. In an A3S Code `use` worker, call +`mcp__use_ocr__ocr_doctor` and `mcp__use_ocr__ocr_extract` directly. The host +owns the MCP process; do not run a shell command, install a provider, or read the +file through another tool. + +## Workflow + +1. Call `mcp__use_ocr__ocr_doctor`. +2. Confirm which provider is ready and whether `sendsSourceOffDevice` is true. +3. Call `mcp__use_ocr__ocr_extract` with the exact local image path from the + task. Supply language identifiers only when known. +4. Preserve the returned source path, media type, size, and SHA-256 in the + result. Treat text, confidence, and bounding boxes as OCR evidence, not as a + verified transcription. + +The local Tesseract provider does not send the image over the network. The +vision provider sends the complete source image and prompt to its disclosed +endpoint. Do not use a non-loopback vision provider unless the user has +authorized that data transfer. Never install, repair, or switch providers from +inside the `use` worker. + +In a CLI-only host, equivalent commands are: + +```bash +a3s use ocr doctor --json +a3s use ocr extract "$IMAGE" --language eng --json +``` + +`a3s-use-ocr` accepts the same arguments when invoked as a standalone +development binary. + +## Boundaries + +- Only bounded local image files are accepted. URLs and PDF rasterization are + outside this domain. +- Keep the default prompt for faithful transcription. A custom vision prompt + must remain an extraction instruction; do not ask the provider to interpret + unrelated content. +- Never report vision output as calibrated confidence or layout evidence. +- Do not silently fall back from a requested provider. Report typed provider, + source, and response errors to the parent agent. diff --git a/crates/ocr/src/cli.rs b/crates/ocr/src/cli.rs new file mode 100644 index 00000000..ea76bcda --- /dev/null +++ b/crates/ocr/src/cli.rs @@ -0,0 +1,205 @@ +use std::path::PathBuf; + +use a3s_use_core::{UseError, UseResult}; +use clap::error::ErrorKind; +use clap::{Parser, Subcommand, ValueEnum}; +use serde::Serialize; + +use crate::{OcrClient, OcrMcpServer, OcrProviderKind, OcrRequest}; + +#[derive(Debug)] +pub struct CommandOutput { + pub human: String, + pub json: serde_json::Value, + pub exit_code: u8, + pub should_print: bool, +} + +impl CommandOutput { + fn data(value: T) -> UseResult + where + T: Serialize, + { + let data = serde_json::to_value(value).map_err(output_error)?; + let human = serde_json::to_string_pretty(&data).map_err(output_error)?; + Ok(Self { + human, + json: serde_json::json!({ + "schemaVersion": 1, + "ok": true, + "data": data, + }), + exit_code: 0, + should_print: true, + }) + } + + fn text(value: String) -> Self { + Self { + human: value.clone(), + json: serde_json::json!({ + "schemaVersion": 1, + "ok": true, + "data": { "text": value }, + }), + exit_code: 0, + should_print: true, + } + } + + fn silent() -> Self { + Self { + human: String::new(), + json: serde_json::Value::Null, + exit_code: 0, + should_print: false, + } + } +} + +#[derive(Debug, Parser)] +#[command( + name = "a3s-use-ocr", + version, + about = "Typed built-in OCR for A3S Use", + arg_required_else_help = true +)] +struct Cli { + /// Emit one versioned JSON document. + #[arg(long, global = true)] + json: bool, + + #[command(subcommand)] + command: Command, +} + +#[derive(Debug, Subcommand)] +enum Command { + /// Inspect provider readiness without reading an image. + Doctor, + /// Extract text and available layout evidence from one local image. + Extract { + path: PathBuf, + /// OCR language identifier; may be repeated. + #[arg(long = "language")] + languages: Vec, + /// Tesseract page segmentation mode from 0 through 13. + #[arg(long = "psm")] + page_segmentation_mode: Option, + /// Override the configured OCR provider for this call. + #[arg(long, value_enum)] + provider: Option, + /// Vision-only extraction instruction. + #[arg(long)] + prompt: Option, + }, + /// Run an extension protocol surface. + Serve { + /// Serve standard MCP over stdin/stdout. + #[arg(long)] + mcp: bool, + }, +} + +#[derive(Debug, Clone, Copy, ValueEnum)] +enum ProviderArg { + Auto, + Tesseract, + Vision, +} + +impl From for OcrProviderKind { + fn from(value: ProviderArg) -> Self { + match value { + ProviderArg::Auto => Self::Auto, + ProviderArg::Tesseract => Self::Tesseract, + ProviderArg::Vision => Self::Vision, + } + } +} + +pub async fn run(args: Vec) -> UseResult { + let mut argv = vec!["a3s-use-ocr".to_string()]; + argv.extend(args); + let cli = match Cli::try_parse_from(argv) { + Ok(cli) => cli, + Err(error) + if matches!( + error.kind(), + ErrorKind::DisplayHelp | ErrorKind::DisplayVersion + ) => + { + return Ok(CommandOutput::text(error.to_string())); + } + Err(error) => return Err(usage_error(error.to_string())), + }; + + if let Command::Serve { mcp } = &cli.command { + if !mcp { + return Err(usage_error("serve requires --mcp")); + } + if cli.json { + return Err(usage_error("--json cannot be combined with serve --mcp")); + } + OcrMcpServer::from_env()?.serve_stdio().await?; + return Ok(CommandOutput::silent()); + } + + let client = OcrClient::from_env()?; + match cli.command { + Command::Doctor => CommandOutput::data(client.diagnostic()), + Command::Extract { + path, + languages, + page_segmentation_mode, + provider, + prompt, + } => CommandOutput::data( + client + .extract(OcrRequest { + path, + languages, + page_segmentation_mode, + provider: provider.map(Into::into), + prompt, + }) + .await?, + ), + Command::Serve { .. } => Err(UseError::new( + "use.ocr.command_invalid", + "OCR MCP command dispatch reached an invalid state.", + )), + } +} + +fn output_error(error: serde_json::Error) -> UseError { + UseError::new( + "use.ocr.output_invalid", + format!("Failed to encode OCR command output: {error}"), + ) +} + +fn usage_error(message: impl Into) -> UseError { + UseError::new("use.ocr.usage_invalid", message).with_suggestion("Run 'a3s use ocr --help'.") +} + +#[cfg(test)] +mod tests { + use super::*; + + #[tokio::test] + async fn doctor_is_versioned_even_when_no_provider_is_ready() { + let output = run(vec!["doctor".to_string(), "--json".to_string()]) + .await + .unwrap(); + assert_eq!(output.json["schemaVersion"], 1); + assert_eq!(output.json["ok"], true); + assert!(output.json["data"]["readiness"].is_string()); + } + + #[tokio::test] + async fn serve_requires_an_explicit_protocol() { + let error = run(vec!["serve".to_string()]).await.unwrap_err(); + assert_eq!(error.code, "use.ocr.usage_invalid"); + } +} diff --git a/crates/ocr/src/client.rs b/crates/ocr/src/client.rs new file mode 100644 index 00000000..21fb3494 --- /dev/null +++ b/crates/ocr/src/client.rs @@ -0,0 +1,649 @@ +use std::collections::BTreeMap; +use std::path::Path; +#[cfg(all(test, unix))] +use std::path::PathBuf; +use std::process::Stdio; + +use a3s_use_core::{Artifact, UseError, UseResult}; +use base64::Engine; +use sha2::{Digest, Sha256}; +use tokio::io::AsyncReadExt; +use tokio::process::Command; + +use crate::models::{OcrBlock, OcrBoundingBox, OcrProviderKind, OcrRequest, OcrResult}; +use crate::provider::{Provider, ProviderConfig}; +use crate::OcrDiagnostic; + +const MAX_INPUT_BYTES: u64 = 32 * 1024 * 1024; +const MAX_PROVIDER_OUTPUT_BYTES: usize = 8 * 1024 * 1024; +const DEFAULT_VISION_PROMPT: &str = "Transcribe all visible text in reading order. Preserve line breaks and meaningful spacing. Return only the transcription; do not summarize, translate, or wrap it in Markdown."; + +#[derive(Clone)] +pub struct OcrClient { + providers: ProviderConfig, + http: reqwest::Client, +} + +impl OcrClient { + pub fn from_env() -> UseResult { + Self::from_provider_config(ProviderConfig::from_env()?) + } + + fn from_provider_config(providers: ProviderConfig) -> UseResult { + let http = reqwest::Client::builder() + .user_agent(concat!("a3s-use-ocr/", env!("CARGO_PKG_VERSION"))) + .build() + .map_err(|error| { + UseError::new( + "use.ocr.client_failed", + format!("Failed to initialize the OCR HTTP client: {error}"), + ) + })?; + Ok(Self { providers, http }) + } + + #[cfg(all(test, unix))] + pub(crate) fn with_tesseract(executable: PathBuf) -> UseResult { + Self::from_provider_config(ProviderConfig::tesseract(executable)) + } + + pub fn diagnostic(&self) -> OcrDiagnostic { + self.providers.diagnostic() + } + + pub async fn extract(&self, request: OcrRequest) -> UseResult { + validate_request(&request)?; + let source = read_source(&request.path).await?; + let provider = self + .providers + .resolve(request.provider.unwrap_or(OcrProviderKind::Auto))?; + let languages = if request.languages.is_empty() { + vec!["eng".to_string()] + } else { + request.languages.clone() + }; + + let (text, blocks, warnings) = match &provider { + Provider::Tesseract { + executable, + timeout, + } => { + let output = run_tesseract( + executable, + &source.artifact.path, + &languages, + request.page_segmentation_mode, + *timeout, + ) + .await?; + let (text, blocks) = parse_tesseract_tsv(&output)?; + (text, blocks, Vec::new()) + } + Provider::Vision { + endpoint, + api_key, + model, + timeout, + } => { + let text = self + .run_vision( + endpoint, + api_key.as_deref(), + model, + &source, + request.prompt.as_deref(), + *timeout, + ) + .await?; + let blocks = (!text.is_empty()) + .then(|| OcrBlock { + page: 1, + text: text.clone(), + confidence: None, + bounding_box: None, + }) + .into_iter() + .collect(); + ( + text, + blocks, + vec![ + "The vision provider does not return calibrated OCR confidence or bounding boxes." + .to_string(), + ], + ) + } + }; + + Ok(OcrResult { + provider: provider.kind(), + source: source.artifact, + languages, + text, + blocks, + warnings, + }) + } + + async fn run_vision( + &self, + endpoint: &url::Url, + api_key: Option<&str>, + model: &str, + source: &SourceImage, + prompt: Option<&str>, + timeout: std::time::Duration, + ) -> UseResult { + let encoded = base64::engine::general_purpose::STANDARD.encode(&source.bytes); + let data_url = format!("data:{};base64,{encoded}", source.artifact.media_type); + let prompt = prompt + .map(str::trim) + .filter(|value| !value.is_empty()) + .unwrap_or(DEFAULT_VISION_PROMPT); + let body = serde_json::json!({ + "model": model, + "temperature": 0, + "messages": [{ + "role": "user", + "content": [ + { "type": "text", "text": prompt }, + { + "type": "image_url", + "image_url": { + "url": data_url, + "detail": "high" + } + } + ] + }] + }); + let mut request = self + .http + .post(endpoint.clone()) + .timeout(timeout) + .json(&body); + if let Some(api_key) = api_key { + request = request.bearer_auth(api_key); + } + let response = request.send().await.map_err(|error| { + UseError::new( + "use.ocr.vision_request_failed", + format!("The vision OCR request failed: {error}"), + ) + .with_detail("endpoint", redacted_endpoint(endpoint)) + })?; + let status = response.status(); + let bytes = response.bytes().await.map_err(|error| { + UseError::new( + "use.ocr.vision_response_invalid", + format!("Failed to read the vision OCR response: {error}"), + ) + })?; + if bytes.len() > MAX_PROVIDER_OUTPUT_BYTES { + return Err(UseError::new( + "use.ocr.output_too_large", + "The vision OCR provider response exceeded 8 MiB.", + )); + } + if !status.is_success() { + let message = String::from_utf8_lossy(&bytes); + return Err(UseError::new( + "use.ocr.vision_request_failed", + format!( + "The vision OCR provider returned HTTP {status}: {}", + bounded_text(&message, 1024) + ), + ) + .with_detail("status", u64::from(status.as_u16()))); + } + let value: serde_json::Value = serde_json::from_slice(&bytes).map_err(|error| { + UseError::new( + "use.ocr.vision_response_invalid", + format!("The vision OCR provider returned invalid JSON: {error}"), + ) + })?; + let content = value.pointer("/choices/0/message/content").ok_or_else(|| { + UseError::new( + "use.ocr.vision_response_invalid", + "The vision OCR response did not contain choices[0].message.content.", + ) + })?; + let text = vision_content_text(content)?; + Ok(text.trim().to_string()) + } +} + +struct SourceImage { + artifact: Artifact, + bytes: Vec, +} + +async fn read_source(path: &Path) -> UseResult { + let canonical = tokio::fs::canonicalize(path).await.map_err(|error| { + UseError::new( + "use.ocr.source_unreadable", + format!("Failed to resolve OCR source '{}': {error}", path.display()), + ) + })?; + let metadata = tokio::fs::metadata(&canonical).await.map_err(|error| { + UseError::new( + "use.ocr.source_unreadable", + format!( + "Failed to inspect OCR source '{}': {error}", + canonical.display() + ), + ) + })?; + if !metadata.is_file() { + return Err(UseError::new( + "use.ocr.source_invalid", + format!( + "OCR source '{}' is not a regular file.", + canonical.display() + ), + )); + } + if metadata.len() == 0 || metadata.len() > MAX_INPUT_BYTES { + return Err(UseError::new( + "use.ocr.source_too_large", + format!( + "OCR source '{}' must contain between 1 byte and 32 MiB.", + canonical.display() + ), + ) + .with_detail("size", metadata.len())); + } + let file = tokio::fs::File::open(&canonical).await.map_err(|error| { + UseError::new( + "use.ocr.source_unreadable", + format!( + "Failed to open OCR source '{}': {error}", + canonical.display() + ), + ) + })?; + let mut bytes = Vec::with_capacity(metadata.len().min(MAX_INPUT_BYTES) as usize); + file.take(MAX_INPUT_BYTES + 1) + .read_to_end(&mut bytes) + .await + .map_err(|error| { + UseError::new( + "use.ocr.source_unreadable", + format!( + "Failed to read OCR source '{}': {error}", + canonical.display() + ), + ) + })?; + if bytes.len() as u64 > MAX_INPUT_BYTES { + return Err(UseError::new( + "use.ocr.source_too_large", + format!( + "OCR source '{}' must not exceed 32 MiB.", + canonical.display() + ), + ) + .with_detail("sizeAtLeast", MAX_INPUT_BYTES + 1)); + } + let media_type = detect_image_type(&bytes).ok_or_else(|| { + UseError::new( + "use.ocr.source_type_unsupported", + "OCR accepts PNG, JPEG, WebP, GIF, BMP, and TIFF image bytes.", + ) + })?; + let digest = Sha256::digest(&bytes); + Ok(SourceImage { + artifact: Artifact { + path: canonical, + media_type: media_type.to_string(), + size: bytes.len() as u64, + sha256: format!("{digest:x}"), + }, + bytes, + }) +} + +fn validate_request(request: &OcrRequest) -> UseResult<()> { + if request.languages.len() > 16 { + return Err(UseError::new( + "use.ocr.languages_invalid", + "At most 16 OCR language identifiers may be requested.", + )); + } + for language in &request.languages { + if language.is_empty() + || language.len() > 32 + || !language + .bytes() + .all(|byte| byte.is_ascii_alphanumeric() || matches!(byte, b'_' | b'-')) + { + return Err(UseError::new( + "use.ocr.languages_invalid", + format!("OCR language identifier '{language}' is invalid."), + )); + } + } + if request.page_segmentation_mode.is_some_and(|mode| mode > 13) { + return Err(UseError::new( + "use.ocr.page_segmentation_invalid", + "Tesseract page segmentation mode must be from 0 through 13.", + )); + } + if request + .prompt + .as_ref() + .is_some_and(|prompt| prompt.len() > 8 * 1024) + { + return Err(UseError::new( + "use.ocr.prompt_too_large", + "The vision OCR prompt must not exceed 8192 bytes.", + )); + } + Ok(()) +} + +async fn run_tesseract( + executable: &Path, + source: &Path, + languages: &[String], + page_segmentation_mode: Option, + timeout: std::time::Duration, +) -> UseResult> { + let mut command = Command::new(executable); + command + .arg(source) + .arg("stdout") + .arg("-l") + .arg(languages.join("+")) + .stdout(Stdio::piped()) + .stderr(Stdio::piped()) + .kill_on_drop(true); + if let Some(mode) = page_segmentation_mode { + command.arg("--psm").arg(mode.to_string()); + } + command.arg("tsv"); + + let output = tokio::time::timeout(timeout, command.output()) + .await + .map_err(|_| { + UseError::new( + "use.ocr.provider_timeout", + format!( + "Tesseract exceeded the {} ms OCR timeout.", + timeout.as_millis() + ), + ) + })? + .map_err(|error| { + UseError::new( + "use.ocr.provider_failed", + format!( + "Failed to launch Tesseract executable '{}': {error}", + executable.display() + ), + ) + })?; + if output.stdout.len() > MAX_PROVIDER_OUTPUT_BYTES + || output.stderr.len() > MAX_PROVIDER_OUTPUT_BYTES + { + return Err(UseError::new( + "use.ocr.output_too_large", + "Tesseract output exceeded 8 MiB.", + )); + } + if !output.status.success() { + let stderr = String::from_utf8_lossy(&output.stderr); + return Err(UseError::new( + "use.ocr.provider_failed", + format!( + "Tesseract exited with {}: {}", + output.status, + bounded_text(&stderr, 2048) + ), + )); + } + Ok(output.stdout) +} + +#[derive(Default)] +struct LineAccumulator { + page: u32, + words: Vec, + confidence_sum: f32, + confidence_count: usize, + left: u32, + top: u32, + right: u32, + bottom: u32, +} + +fn parse_tesseract_tsv(output: &[u8]) -> UseResult<(String, Vec)> { + let output = std::str::from_utf8(output).map_err(|error| { + UseError::new( + "use.ocr.provider_output_invalid", + format!("Tesseract TSV output was not UTF-8: {error}"), + ) + })?; + let mut lines = BTreeMap::<(u32, u32, u32, u32), LineAccumulator>::new(); + for (index, row) in output.lines().enumerate() { + if index == 0 && row.starts_with("level\t") { + continue; + } + if row.trim().is_empty() { + continue; + } + let columns = row.splitn(12, '\t').collect::>(); + if columns.len() != 12 { + return Err(UseError::new( + "use.ocr.provider_output_invalid", + format!( + "Tesseract TSV row {} did not contain 12 columns.", + index + 1 + ), + )); + } + let level = parse_u32(columns[0], index)?; + if level != 5 || columns[11].trim().is_empty() { + continue; + } + let page = parse_u32(columns[1], index)?; + let block = parse_u32(columns[2], index)?; + let paragraph = parse_u32(columns[3], index)?; + let line = parse_u32(columns[4], index)?; + let left = parse_u32(columns[6], index)?; + let top = parse_u32(columns[7], index)?; + let width = parse_u32(columns[8], index)?; + let height = parse_u32(columns[9], index)?; + let confidence = columns[10] + .parse::() + .ok() + .filter(|value| *value >= 0.0); + let entry = lines + .entry((page, block, paragraph, line)) + .or_insert_with(|| LineAccumulator { + page, + left, + top, + right: left.saturating_add(width), + bottom: top.saturating_add(height), + ..LineAccumulator::default() + }); + entry.words.push(columns[11].trim().to_string()); + if let Some(confidence) = confidence { + entry.confidence_sum += confidence; + entry.confidence_count += 1; + } + entry.left = entry.left.min(left); + entry.top = entry.top.min(top); + entry.right = entry.right.max(left.saturating_add(width)); + entry.bottom = entry.bottom.max(top.saturating_add(height)); + } + let blocks = lines + .into_values() + .filter_map(|line| { + let text = line.words.join(" "); + (!text.is_empty()).then(|| OcrBlock { + page: line.page, + text, + confidence: (line.confidence_count > 0) + .then(|| line.confidence_sum / line.confidence_count as f32), + bounding_box: Some(OcrBoundingBox { + x: line.left, + y: line.top, + width: line.right.saturating_sub(line.left), + height: line.bottom.saturating_sub(line.top), + }), + }) + }) + .collect::>(); + let text = blocks + .iter() + .map(|block| block.text.as_str()) + .collect::>() + .join("\n"); + Ok((text, blocks)) +} + +fn parse_u32(value: &str, row: usize) -> UseResult { + value.parse::().map_err(|_| { + UseError::new( + "use.ocr.provider_output_invalid", + format!( + "Tesseract TSV row {} contained an invalid integer.", + row + 1 + ), + ) + }) +} + +fn vision_content_text(content: &serde_json::Value) -> UseResult { + if let Some(text) = content.as_str() { + return Ok(text.to_string()); + } + let Some(parts) = content.as_array() else { + return Err(UseError::new( + "use.ocr.vision_response_invalid", + "Vision OCR message content was neither text nor a text-part array.", + )); + }; + let text = parts + .iter() + .filter_map(|part| { + part.get("text") + .and_then(serde_json::Value::as_str) + .or_else(|| part.as_str()) + }) + .collect::>() + .join(""); + if text.is_empty() { + return Err(UseError::new( + "use.ocr.vision_response_invalid", + "Vision OCR message content did not contain text.", + )); + } + Ok(text) +} + +fn detect_image_type(bytes: &[u8]) -> Option<&'static str> { + if bytes.starts_with(b"\x89PNG\r\n\x1a\n") { + Some("image/png") + } else if bytes.starts_with(b"\xff\xd8\xff") { + Some("image/jpeg") + } else if bytes.starts_with(b"GIF87a") || bytes.starts_with(b"GIF89a") { + Some("image/gif") + } else if bytes.starts_with(b"BM") { + Some("image/bmp") + } else if bytes.starts_with(b"II*\0") || bytes.starts_with(b"MM\0*") { + Some("image/tiff") + } else if bytes.len() >= 12 && bytes.starts_with(b"RIFF") && &bytes[8..12] == b"WEBP" { + Some("image/webp") + } else { + None + } +} + +fn bounded_text(value: &str, max: usize) -> String { + let mut text = value.chars().take(max).collect::(); + if value.chars().count() > max { + text.push('…'); + } + text +} + +fn redacted_endpoint(endpoint: &url::Url) -> String { + let mut redacted = endpoint.clone(); + redacted.set_query(None); + redacted.set_fragment(None); + redacted.to_string() +} + +#[cfg(test)] +mod tests { + use super::*; + + #[cfg(unix)] + use std::os::unix::fs::PermissionsExt; + + #[test] + fn parses_tesseract_words_into_ordered_lines() { + let tsv = b"level\tpage_num\tblock_num\tpar_num\tline_num\tword_num\tleft\ttop\twidth\theight\tconf\ttext\n5\t1\t1\t1\t1\t1\t10\t20\t30\t10\t95.0\tHello\n5\t1\t1\t1\t1\t2\t45\t20\t35\t10\t85.0\tworld\n5\t1\t1\t1\t2\t1\t10\t40\t20\t10\t90.0\tNext\n"; + let (text, blocks) = parse_tesseract_tsv(tsv).unwrap(); + assert_eq!(text, "Hello world\nNext"); + assert_eq!(blocks.len(), 2); + assert_eq!(blocks[0].confidence, Some(90.0)); + assert_eq!( + blocks[0].bounding_box, + Some(OcrBoundingBox { + x: 10, + y: 20, + width: 70, + height: 10, + }) + ); + } + + #[test] + fn detects_supported_image_signatures() { + assert_eq!( + detect_image_type(b"\x89PNG\r\n\x1a\nrest"), + Some("image/png") + ); + assert_eq!(detect_image_type(b"\xff\xd8\xffrest"), Some("image/jpeg")); + assert_eq!(detect_image_type(b"not an image"), None); + } + + #[cfg(unix)] + #[tokio::test] + async fn local_provider_extracts_a_real_bounded_source_through_its_process_boundary() { + let temp = tempfile::tempdir().unwrap(); + let executable = temp.path().join("tesseract-fixture"); + std::fs::write( + &executable, + "#!/bin/sh\nprintf 'level\\tpage_num\\tblock_num\\tpar_num\\tline_num\\tword_num\\tleft\\ttop\\twidth\\theight\\tconf\\ttext\\n5\\t1\\t1\\t1\\t1\\t1\\t2\\t3\\t20\\t8\\t98.0\\tA3S\\n5\\t1\\t1\\t1\\t1\\t2\\t24\\t3\\t30\\t8\\t96.0\\tUse\\n'\n", + ) + .unwrap(); + let mut permissions = std::fs::metadata(&executable).unwrap().permissions(); + permissions.set_mode(0o755); + std::fs::set_permissions(&executable, permissions).unwrap(); + let image = temp.path().join("scan.png"); + std::fs::write(&image, b"\x89PNG\r\n\x1a\nfixture").unwrap(); + + let result = OcrClient::with_tesseract(executable) + .unwrap() + .extract(OcrRequest { + path: image, + languages: vec!["eng".to_string()], + page_segmentation_mode: Some(6), + provider: Some(OcrProviderKind::Tesseract), + prompt: None, + }) + .await + .unwrap(); + + assert_eq!(result.provider, OcrProviderKind::Tesseract); + assert_eq!(result.text, "A3S Use"); + assert_eq!(result.blocks.len(), 1); + assert_eq!(result.source.media_type, "image/png"); + assert_eq!(result.source.sha256.len(), 64); + } +} diff --git a/crates/ocr/src/lib.rs b/crates/ocr/src/lib.rs new file mode 100644 index 00000000..a1bd1740 --- /dev/null +++ b/crates/ocr/src/lib.rs @@ -0,0 +1,18 @@ +//! Typed optical character recognition for A3S Use. +//! +//! OCR is a first-party built-in Use domain and remains process-isolated from +//! A3S Code through its standard MCP server. The crate supports a local +//! Tesseract executable and an explicitly configured OpenAI-compatible vision +//! endpoint without silently installing either provider. + +pub mod cli; +mod client; +pub mod mcp; +mod models; +mod provider; + +pub use client::OcrClient; +pub use mcp::OcrMcpServer; +pub use models::{OcrBlock, OcrBoundingBox, OcrDiagnostic, OcrProviderKind, OcrRequest, OcrResult}; + +pub use a3s_use_core::{Artifact, Readiness, UseError, UseResult}; diff --git a/crates/ocr/src/main.rs b/crates/ocr/src/main.rs new file mode 100644 index 00000000..4bba3c58 --- /dev/null +++ b/crates/ocr/src/main.rs @@ -0,0 +1,39 @@ +use std::process::ExitCode; + +#[tokio::main] +async fn main() -> ExitCode { + let args = std::env::args().skip(1).collect::>(); + let json = args.iter().any(|argument| argument == "--json"); + match a3s_use_ocr::cli::run(args).await { + Ok(output) => { + if output.should_print && json { + println!( + "{}", + serde_json::to_string_pretty(&output.json).unwrap_or_default() + ); + } else if output.should_print && !output.human.is_empty() { + println!("{}", output.human); + } + ExitCode::from(output.exit_code) + } + Err(error) => { + if json { + let output = serde_json::json!({ + "schemaVersion": 1, + "ok": false, + "error": error, + }); + println!( + "{}", + serde_json::to_string_pretty(&output).unwrap_or_default() + ); + } else { + eprintln!("a3s-use-ocr: {error}"); + if let Some(suggestion) = &error.suggestion { + eprintln!("suggestion: {suggestion}"); + } + } + ExitCode::from(1) + } + } +} diff --git a/crates/ocr/src/mcp.rs b/crates/ocr/src/mcp.rs new file mode 100644 index 00000000..f2de4ada --- /dev/null +++ b/crates/ocr/src/mcp.rs @@ -0,0 +1,170 @@ +//! Standard MCP tools for the built-in OCR domain. + +use rmcp::handler::server::{router::tool::ToolRouter, wrapper::Parameters}; +use rmcp::model::{CallToolResult, Implementation, ServerCapabilities, ServerInfo}; +use rmcp::{tool, tool_handler, tool_router, ServerHandler, ServiceExt}; +use serde::Serialize; + +use crate::{OcrClient, OcrDiagnostic, OcrRequest, OcrResult, UseError, UseResult}; + +#[derive(Clone)] +pub struct OcrMcpServer { + client: OcrClient, + tool_router: ToolRouter, +} + +impl OcrMcpServer { + pub fn new(client: OcrClient) -> Self { + Self { + client, + tool_router: Self::tool_router(), + } + } + + pub fn from_env() -> UseResult { + Ok(Self::new(OcrClient::from_env()?)) + } + + /// Serve standard MCP framing over stdin/stdout until the peer disconnects. + pub async fn serve_stdio(self) -> UseResult<()> { + let service = self + .serve(rmcp::transport::stdio()) + .await + .map_err(|error| mcp_error("start", error))?; + service + .waiting() + .await + .map_err(|error| mcp_error("run", error))?; + Ok(()) + } +} + +#[tool_router] +impl OcrMcpServer { + #[tool( + name = "ocr_doctor", + description = "Inspect OCR provider readiness without reading an image or making a network request", + output_schema = rmcp::handler::server::tool::cached_schema_for_type::(), + annotations( + read_only_hint = true, + destructive_hint = false, + idempotent_hint = true, + open_world_hint = false + ) + )] + async fn ocr_doctor(&self) -> Result { + Ok(tool_result(Ok(self.client.diagnostic()))) + } + + #[tool( + name = "ocr_extract", + description = "Extract text and available layout evidence from one bounded local image; the configured vision provider may send source bytes to its disclosed endpoint", + output_schema = rmcp::handler::server::tool::cached_schema_for_type::(), + annotations( + read_only_hint = true, + destructive_hint = false, + idempotent_hint = true, + open_world_hint = true + ) + )] + async fn ocr_extract( + &self, + Parameters(request): Parameters, + ) -> Result { + Ok(tool_result(self.client.extract(request).await)) + } +} + +#[tool_handler] +impl ServerHandler for OcrMcpServer { + fn get_info(&self) -> ServerInfo { + ServerInfo { + capabilities: ServerCapabilities::builder().enable_tools().build(), + server_info: Implementation { + name: "a3s-use-ocr".to_string(), + title: Some("A3S Use OCR".to_string()), + version: env!("CARGO_PKG_VERSION").to_string(), + icons: None, + website_url: Some("https://github.com/A3S-Lab/Use".to_string()), + }, + instructions: Some( + "Call ocr_doctor first. Use ocr_extract only for a local image path supplied by the task. A vision provider may send the complete image to its configured endpoint; do not use it without the user's authority. Preserve the source SHA-256 and distinguish OCR text from verified source text." + .to_string(), + ), + ..Default::default() + } + } +} + +fn tool_result(result: UseResult) -> CallToolResult +where + T: Serialize, +{ + match result { + Ok(output) => match serde_json::to_value(output) { + Ok(value) => CallToolResult::structured(value), + Err(error) => tool_error(UseError::new( + "use.ocr.output_invalid", + format!("Failed to encode OCR MCP output: {error}"), + )), + }, + Err(error) => tool_error(error), + } +} + +fn tool_error(error: UseError) -> CallToolResult { + CallToolResult::structured_error(serde_json::to_value(error).unwrap_or_else(|_| { + serde_json::json!({ + "code": "use.error_encoding_failed", + "message": "Failed to encode A3S Use error." + }) + })) +} + +fn mcp_error(action: &str, error: impl std::fmt::Display) -> UseError { + UseError::new( + "use.ocr.mcp_failed", + format!("Failed to {action} the OCR MCP server: {error}"), + ) +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn server_exposes_typed_annotated_ocr_tools() { + let client = OcrClient::from_env().unwrap(); + let server = OcrMcpServer::new(client); + let mut tools = server.tool_router.list_all(); + tools.sort_by(|left, right| left.name.cmp(&right.name)); + assert_eq!( + tools + .iter() + .map(|tool| tool.name.as_ref()) + .collect::>(), + ["ocr_doctor", "ocr_extract"] + ); + let doctor = tools.iter().find(|tool| tool.name == "ocr_doctor").unwrap(); + let extract = tools + .iter() + .find(|tool| tool.name == "ocr_extract") + .unwrap(); + assert!(doctor.output_schema.is_some()); + assert!(extract.output_schema.is_some()); + assert_eq!( + doctor + .annotations + .as_ref() + .and_then(|annotations| annotations.open_world_hint), + Some(false) + ); + assert_eq!( + extract + .annotations + .as_ref() + .and_then(|annotations| annotations.open_world_hint), + Some(true) + ); + } +} diff --git a/crates/ocr/src/models.rs b/crates/ocr/src/models.rs new file mode 100644 index 00000000..6cbf6bae --- /dev/null +++ b/crates/ocr/src/models.rs @@ -0,0 +1,105 @@ +use std::path::PathBuf; + +use a3s_use_core::{Artifact, Readiness}; +use serde::{Deserialize, Serialize}; + +#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize, schemars::JsonSchema)] +#[serde(rename_all = "kebab-case")] +pub enum OcrProviderKind { + Auto, + Tesseract, + Vision, +} + +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize, schemars::JsonSchema)] +#[serde(rename_all = "camelCase")] +pub struct OcrRequest { + #[schemars(description = "Local PNG, JPEG, WebP, GIF, BMP, or TIFF image path")] + pub path: PathBuf, + #[serde(default)] + #[schemars( + description = "OCR language identifiers; Tesseract values are joined with '+', for example ['eng', 'chi_sim']" + )] + pub languages: Vec, + #[serde(default)] + #[schemars(description = "Optional Tesseract page segmentation mode from 0 through 13")] + pub page_segmentation_mode: Option, + #[serde(default)] + #[schemars(description = "Override the configured provider for this call")] + pub provider: Option, + #[serde(default)] + #[schemars(description = "Optional extraction instruction used only by the vision provider")] + pub prompt: Option, +} + +#[derive(Debug, Clone, PartialEq, Serialize, Deserialize, schemars::JsonSchema)] +#[serde(rename_all = "camelCase")] +pub struct OcrBoundingBox { + pub x: u32, + pub y: u32, + pub width: u32, + pub height: u32, +} + +#[derive(Debug, Clone, PartialEq, Serialize, Deserialize, schemars::JsonSchema)] +#[serde(rename_all = "camelCase")] +pub struct OcrBlock { + pub page: u32, + pub text: String, + #[serde(skip_serializing_if = "Option::is_none")] + pub confidence: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub bounding_box: Option, +} + +#[derive(Debug, Clone, PartialEq, Serialize, Deserialize, schemars::JsonSchema)] +#[serde(rename_all = "camelCase")] +pub struct OcrResult { + pub provider: OcrProviderKind, + #[schemars(with = "OcrArtifactSchema")] + pub source: Artifact, + pub languages: Vec, + pub text: String, + #[serde(default, skip_serializing_if = "Vec::is_empty")] + pub blocks: Vec, + #[serde(default, skip_serializing_if = "Vec::is_empty")] + pub warnings: Vec, +} + +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize, schemars::JsonSchema)] +#[serde(rename_all = "camelCase")] +pub struct OcrDiagnostic { + #[schemars(with = "OcrReadinessSchema")] + pub readiness: Readiness, + #[serde(skip_serializing_if = "Option::is_none")] + pub provider: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub executable: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub endpoint: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub model: Option, + pub sends_source_off_device: bool, + pub message: String, + #[serde(default, skip_serializing_if = "Vec::is_empty")] + pub suggestions: Vec, +} + +#[derive(schemars::JsonSchema)] +#[allow(dead_code)] +struct OcrArtifactSchema { + path: PathBuf, + media_type: String, + size: u64, + sha256: String, +} + +#[derive(schemars::JsonSchema)] +#[serde(rename_all = "kebab-case")] +#[allow(dead_code)] +enum OcrReadinessSchema { + Ready, + Missing, + Broken, + Unknown, +} diff --git a/crates/ocr/src/provider.rs b/crates/ocr/src/provider.rs new file mode 100644 index 00000000..eb8880e0 --- /dev/null +++ b/crates/ocr/src/provider.rs @@ -0,0 +1,393 @@ +use std::env; +use std::path::{Path, PathBuf}; +use std::time::Duration; + +use a3s_use_core::{Readiness, UseError, UseResult}; +use url::Url; + +use crate::{OcrDiagnostic, OcrProviderKind}; + +const DEFAULT_TIMEOUT: Duration = Duration::from_secs(60); +const DEFAULT_VISION_BASE_URL: &str = "https://api.openai.com/v1/"; + +#[derive(Debug, Clone)] +pub(crate) enum Provider { + Tesseract { + executable: PathBuf, + timeout: Duration, + }, + Vision { + endpoint: Url, + api_key: Option, + model: String, + timeout: Duration, + }, +} + +impl Provider { + pub(crate) fn kind(&self) -> OcrProviderKind { + match self { + Self::Tesseract { .. } => OcrProviderKind::Tesseract, + Self::Vision { .. } => OcrProviderKind::Vision, + } + } + + pub(crate) fn diagnostic(&self) -> OcrDiagnostic { + match self { + Self::Tesseract { executable, .. } => OcrDiagnostic { + readiness: Readiness::Ready, + provider: Some(OcrProviderKind::Tesseract), + executable: Some(executable.clone()), + endpoint: None, + model: None, + sends_source_off_device: false, + message: "The local Tesseract OCR provider is ready.".to_string(), + suggestions: Vec::new(), + }, + Self::Vision { + endpoint, model, .. + } => OcrDiagnostic { + readiness: Readiness::Ready, + provider: Some(OcrProviderKind::Vision), + executable: None, + endpoint: Some(redacted_endpoint(endpoint)), + model: Some(model.clone()), + sends_source_off_device: !is_loopback(endpoint), + message: "The explicitly configured vision OCR provider is ready.".to_string(), + suggestions: Vec::new(), + }, + } + } +} + +#[derive(Debug, Clone)] +pub(crate) struct ProviderConfig { + requested: OcrProviderKind, + tesseract: Option, + vision: Option, + timeout: Duration, +} + +#[derive(Debug, Clone)] +struct VisionConfig { + endpoint: Url, + api_key: Option, + model: String, +} + +impl ProviderConfig { + pub(crate) fn from_env() -> UseResult { + let requested = match env::var("A3S_OCR_PROVIDER") + .unwrap_or_else(|_| "auto".to_string()) + .trim() + { + "" | "auto" => OcrProviderKind::Auto, + "tesseract" => OcrProviderKind::Tesseract, + "vision" => OcrProviderKind::Vision, + value => { + return Err(UseError::new( + "use.ocr.provider_invalid", + format!("Unknown OCR provider '{value}'; expected auto, tesseract, or vision."), + )) + } + }; + + let timeout = timeout_from_env()?; + let tesseract = env::var_os("A3S_OCR_TESSERACT_EXECUTABLE") + .filter(|value| !value.is_empty()) + .map(PathBuf::from) + .or_else(|| find_on_path("tesseract")); + let vision = vision_config_from_env()?; + Ok(Self { + requested, + tesseract, + vision, + timeout, + }) + } + + #[cfg(all(test, unix))] + pub(crate) fn tesseract(executable: PathBuf) -> Self { + Self { + requested: OcrProviderKind::Tesseract, + tesseract: Some(executable), + vision: None, + timeout: DEFAULT_TIMEOUT, + } + } + + pub(crate) fn diagnostic(&self) -> OcrDiagnostic { + match self.resolve(self.requested) { + Ok(provider) => provider.diagnostic(), + Err(error) => OcrDiagnostic { + readiness: Readiness::Missing, + provider: match self.requested { + OcrProviderKind::Auto => None, + provider => Some(provider), + }, + executable: self.tesseract.clone(), + endpoint: self + .vision + .as_ref() + .map(|vision| redacted_endpoint(&vision.endpoint)), + model: self.vision.as_ref().map(|vision| vision.model.clone()), + sends_source_off_device: self + .vision + .as_ref() + .is_some_and(|vision| !is_loopback(&vision.endpoint)), + message: error.message, + suggestions: error.suggestion.into_iter().collect(), + }, + } + } + + pub(crate) fn resolve(&self, requested: OcrProviderKind) -> UseResult { + let requested = if requested == OcrProviderKind::Auto { + self.requested + } else { + requested + }; + match requested { + OcrProviderKind::Auto => { + if let Some(executable) = &self.tesseract { + return tesseract_provider(executable, self.timeout); + } + if let Some(vision) = &self.vision { + return Ok(vision_provider(vision, self.timeout)); + } + Err(missing_provider()) + } + OcrProviderKind::Tesseract => self + .tesseract + .as_ref() + .ok_or_else(missing_tesseract) + .and_then(|path| tesseract_provider(path, self.timeout)), + OcrProviderKind::Vision => self + .vision + .as_ref() + .map(|vision| vision_provider(vision, self.timeout)) + .ok_or_else(missing_vision), + } + } +} + +fn tesseract_provider(path: &Path, timeout: Duration) -> UseResult { + let path = std::fs::canonicalize(path).map_err(|error| { + UseError::new( + "use.ocr.provider_missing", + format!( + "Configured Tesseract executable '{}' is not readable: {error}", + path.display() + ), + ) + .with_suggestion( + "Install Tesseract explicitly or configure the vision provider; A3S Use will not install an OCR provider automatically.", + ) + })?; + let metadata = std::fs::metadata(&path).map_err(|error| { + UseError::new( + "use.ocr.provider_missing", + format!( + "Configured Tesseract executable '{}' is not readable: {error}", + path.display() + ), + ) + })?; + if !metadata.is_file() { + return Err(UseError::new( + "use.ocr.provider_invalid", + format!( + "Configured Tesseract path '{}' is not a regular file.", + path.display() + ), + )); + } + #[cfg(unix)] + { + use std::os::unix::fs::PermissionsExt; + if metadata.permissions().mode() & 0o111 == 0 { + return Err(UseError::new( + "use.ocr.provider_invalid", + format!( + "Configured Tesseract path '{}' is not executable.", + path.display() + ), + )); + } + } + Ok(Provider::Tesseract { + executable: path, + timeout, + }) +} + +fn vision_provider(config: &VisionConfig, timeout: Duration) -> Provider { + Provider::Vision { + endpoint: config.endpoint.clone(), + api_key: config.api_key.clone(), + model: config.model.clone(), + timeout, + } +} + +fn vision_config_from_env() -> UseResult> { + let model = env::var("A3S_OCR_VISION_MODEL") + .ok() + .map(|value| value.trim().to_string()) + .filter(|value| !value.is_empty()); + let base_url = env::var("A3S_OCR_VISION_BASE_URL") + .ok() + .map(|value| value.trim().to_string()) + .filter(|value| !value.is_empty()); + let api_key = env::var("A3S_OCR_VISION_API_KEY") + .ok() + .map(|value| value.trim().to_string()) + .filter(|value| !value.is_empty()); + + if model.is_none() && base_url.is_none() && api_key.is_none() { + return Ok(None); + } + let model = model.ok_or_else(|| { + UseError::new( + "use.ocr.vision_config_invalid", + "A3S_OCR_VISION_MODEL is required when the vision OCR provider is configured.", + ) + })?; + let mut base = base_url.unwrap_or_else(|| DEFAULT_VISION_BASE_URL.to_string()); + if !base.ends_with('/') { + base.push('/'); + } + let base = Url::parse(&base).map_err(|error| { + UseError::new( + "use.ocr.vision_config_invalid", + format!("A3S_OCR_VISION_BASE_URL is invalid: {error}"), + ) + })?; + validate_endpoint(&base, api_key.as_deref())?; + let endpoint = base.join("chat/completions").map_err(|error| { + UseError::new( + "use.ocr.vision_config_invalid", + format!("Failed to resolve the vision OCR endpoint: {error}"), + ) + })?; + Ok(Some(VisionConfig { + endpoint, + api_key, + model, + })) +} + +fn validate_endpoint(endpoint: &Url, api_key: Option<&str>) -> UseResult<()> { + if !endpoint.username().is_empty() || endpoint.password().is_some() { + return Err(UseError::new( + "use.ocr.vision_config_invalid", + "The vision OCR endpoint must not contain embedded credentials.", + )); + } + if endpoint.scheme() != "https" && !(endpoint.scheme() == "http" && is_loopback(endpoint)) { + return Err(UseError::new( + "use.ocr.vision_config_invalid", + "The vision OCR endpoint must use HTTPS; loopback HTTP is allowed for local providers.", + )); + } + if !is_loopback(endpoint) && api_key.is_none() { + return Err(UseError::new( + "use.ocr.vision_config_invalid", + "A3S_OCR_VISION_API_KEY is required for a non-loopback vision endpoint.", + )); + } + Ok(()) +} + +fn timeout_from_env() -> UseResult { + let Some(value) = env::var("A3S_OCR_TIMEOUT_MS").ok() else { + return Ok(DEFAULT_TIMEOUT); + }; + let millis = value.parse::().map_err(|_| { + UseError::new( + "use.ocr.timeout_invalid", + "A3S_OCR_TIMEOUT_MS must be an integer from 1 through 300000.", + ) + })?; + if !(1..=300_000).contains(&millis) { + return Err(UseError::new( + "use.ocr.timeout_invalid", + "A3S_OCR_TIMEOUT_MS must be an integer from 1 through 300000.", + )); + } + Ok(Duration::from_millis(millis)) +} + +fn find_on_path(name: &str) -> Option { + let path = env::var_os("PATH")?; + env::split_paths(&path) + .map(|directory| directory.join(executable_name(name))) + .find(|candidate| candidate.is_file()) +} + +fn executable_name(name: &str) -> String { + if cfg!(windows) { + format!("{name}.exe") + } else { + name.to_string() + } +} + +fn missing_provider() -> UseError { + UseError::new( + "use.ocr.provider_missing", + "No OCR provider is configured or discoverable.", + ) + .with_suggestion( + "Install Tesseract explicitly, set A3S_OCR_TESSERACT_EXECUTABLE, or configure A3S_OCR_VISION_MODEL, A3S_OCR_VISION_BASE_URL, and A3S_OCR_VISION_API_KEY.", + ) +} + +fn missing_tesseract() -> UseError { + UseError::new( + "use.ocr.provider_missing", + "The Tesseract OCR provider is not installed or configured.", + ) + .with_suggestion( + "Install Tesseract explicitly or set A3S_OCR_TESSERACT_EXECUTABLE; A3S Use will not install it automatically.", + ) +} + +fn missing_vision() -> UseError { + UseError::new( + "use.ocr.provider_missing", + "The vision OCR provider is not configured.", + ) + .with_suggestion( + "Set A3S_OCR_VISION_MODEL and an approved HTTPS endpoint/API key before sending source images to a vision provider.", + ) +} + +fn is_loopback(url: &Url) -> bool { + matches!(url.host_str(), Some("localhost" | "127.0.0.1" | "::1")) +} + +fn redacted_endpoint(url: &Url) -> String { + let mut redacted = url.clone(); + redacted.set_query(None); + redacted.set_fragment(None); + redacted.to_string() +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn rejects_insecure_remote_vision_endpoint() { + let endpoint = Url::parse("http://ocr.example.com/v1/").unwrap(); + let error = validate_endpoint(&endpoint, Some("secret")).unwrap_err(); + assert_eq!(error.code, "use.ocr.vision_config_invalid"); + } + + #[test] + fn permits_loopback_http_without_an_api_key() { + let endpoint = Url::parse("http://127.0.0.1:8080/v1/").unwrap(); + validate_endpoint(&endpoint, None).unwrap(); + } +} diff --git a/crates/office/skills/a3s-use-office/SKILL.md b/crates/office/skills/a3s-use-office/SKILL.md index 6fe1056f..a608ae12 100644 --- a/crates/office/skills/a3s-use-office/SKILL.md +++ b/crates/office/skills/a3s-use-office/SKILL.md @@ -9,10 +9,27 @@ Use A3S Use as the application boundary for Office documents. Prefer the in-process native engine and its typed operations. Use the compatibility route only when the requested operation is not yet native. +Use the host surface that is already available: + +- In an A3S Code `use` worker, call the available + `mcp__use_office__*` tools directly. The host has already started the native + MCP server and owns its lifecycle; do not run shell commands or install a + provider. +- If a requested operation is absent from the native tools, use an available + `mcp__use_office_compat__*` tool only as an explicit compatibility fallback. + If that route is absent, report the missing capability instead of installing, + repairing, or falling back to a shell. +- In a CLI-only agent host, use the `a3s use office native ...` commands below. + ## Workflow 1. Identify and inspect the document before changing it. + In an A3S Code `use` worker, begin with + `mcp__use_office__office_validate`, then open a session and use + `mcp__use_office__office_view`, `office_get`, or `office_query` as needed. + In a CLI-only host, use: + ```bash a3s use office native validate "$FILE" --json a3s use office native view "$FILE" annotated --limit 200 --json @@ -20,7 +37,9 @@ only when the requested operation is not yet native. a3s use office native view "$FILE" issues --json ``` -2. Load the format reference relevant to the task: +2. In a CLI-only agent host, load the format reference relevant to the task. + An A3S Code `use` worker cannot read Skill reference files; rely on this + guidance and the available MCP tool schemas instead. - Read [references/word.md](references/word.md) for `.docx`. - Read [references/spreadsheet.md](references/spreadsheet.md) for `.xlsx`. @@ -52,6 +71,10 @@ when an agent must bound its lifetime. ## Choose the Surface +- In an A3S Code `use` worker, use `mcp__use_office__*` and keep the returned + Office session ID stable until the document is saved and closed. +- Use `mcp__use_office_compat__*` only when the native route lacks the requested + operation and the compatibility tools are actually present. - Use `a3s use office native ... --json` for local automation and scripts. - Use `a3s use mcp serve office-native` for typed, stateful agent sessions. Read [references/mcp.md](references/mcp.md) before using its session tools. diff --git a/crates/office/skills/a3s-use-office/references/mcp.md b/crates/office/skills/a3s-use-office/references/mcp.md index 5d1da61b..cb78ce1e 100644 --- a/crates/office/skills/a3s-use-office/references/mcp.md +++ b/crates/office/skills/a3s-use-office/references/mcp.md @@ -8,7 +8,11 @@ ## Session Workflow -Start the explicit native standard MCP server: +In an A3S Code `use` worker, the host has already started the native server. +Call the available `mcp__use_office__office_*` tools; do not start a process or +run a shell command. Tool names below omit the host prefix for readability. + +In a CLI-only MCP host, start the explicit native standard MCP server: ```bash a3s use mcp serve office-native @@ -504,6 +508,9 @@ Annotated reads include unsaved mutations in the current typed session. Screenshot output requires a no-clobber `.png` path and a ready A3S Browser provider; other native Office tools do not require Browser or OfficeCLI. -Use `a3s use mcp serve office` only for the pinned OfficeCLI compatibility -server. It is a separate standard MCP target and is not the native session -engine. +In an A3S Code `use` worker, use an available +`mcp__use_office_compat__*` tool only when the native vocabulary lacks the +requested operation. In a CLI-only MCP host, `a3s use mcp serve office-compat` +starts the pinned OfficeCLI compatibility server; the legacy +`a3s use mcp serve office` alias remains supported. It is a separate standard +MCP target and is not the native session engine. diff --git a/docs/architecture.md b/docs/architecture.md index c152127b..c165d7b5 100644 --- a/docs/architecture.md +++ b/docs/architecture.md @@ -37,6 +37,15 @@ The package manifest is a3s-use-extension.acl and is parsed by a3s-acl. A3S Use owns identity, routes, trust, activation, and lifecycle around the surfaces. It does not define JSON-RPC methods or convert surfaces implicitly. +`a3s-use-ocr` implements the reserved first-party `ocr` route in the default +Use build. The release packages its content-bound Skill and exposes the native +CLI plus standard stdio MCP without a separate extension install. The process +accepts bounded local image files and binds every result to the canonical +source digest. It uses only an explicitly present Tesseract executable or +configured vision endpoint; it never installs a provider silently. The vision +diagnostic discloses off-device image transfer, and its MCP extraction tool +carries conservative open-world annotations. + ## Hot-plug registry Extension code remains behind native process boundaries. The registry is a @@ -64,8 +73,8 @@ custom RPC protocol, `dlopen`, or restart is required. ### Unified capability projection Resident Code hosts do not need separate discovery paths for built-in and -external domains. `capability snapshot` projects Browser, Office, Box, and -enabled extensions through one schema while preserving each binding's +external domains. `capability snapshot` projects Browser, native Office, OCR, +Box, and enabled extensions through one schema while preserving each binding's `built-in` or `extension` origin. `capability watch` accepts both the extension generation and a content revision. The generation advances for extension lifecycle commits; the SHA-256 revision also detects built-in provider @@ -73,10 +82,15 @@ readiness and packaged Skill changes when the extension generation remains unchanged. Each Skill projection includes an absolute package path and its own lowercase SHA-256, allowing a resident host to reject raced or modified bytes before replacing its live Skill. -The default distribution projects both the Browser Skill and the first-party -`a3s-use-office` Skill. `office skills list|get|path` exposes the latter as -bounded local CLI reads; it never launches the OfficeCLI compatibility -provider. +The default distribution projects the Browser, first-party `a3s-use-office`, +and first-party `a3s-use-ocr` Skills. `office skills list|get|path` exposes the +Office Skill as bounded local CLI reads; it never launches the OfficeCLI +compatibility provider. For resident hosts, `use/office` targets the built-in +`office-native` MCP server and is ready independently of OfficeCLI. A discovered +OfficeCLI provider is projected separately as `use/office-compat`, targeting +the standard compatibility server without carrying the native Skill. The +`use/ocr` route targets `ocr-native`; provider readiness remains explicit and +never triggers a silent Tesseract or vision-provider install. The projection contains content-bound Skill references and an MCP launch target, never executable extension code or a generic action payload. Consumers still @@ -361,11 +375,13 @@ merged-span rewriting fail before save; Presentation table merge editing remains outside this bounded milestone. These mutations use the existing typed batch transaction and do not introduce another protocol or runtime. -Unpromoted commands are delegated to OfficeCLI and `mcp serve office` launches -its standard MCP server. That compatibility process remains isolated from the +Unpromoted commands are delegated to OfficeCLI and +`mcp serve office-compat` launches its standard MCP server; `mcp serve office` +remains a legacy alias. That compatibility process remains isolated from the native engine. `mcp serve office-native` instead runs the A3S-owned server in -process, never discovers or starts OfficeCLI, and keeps the compatibility target -unchanged until the native product gates pass. +process and never discovers or starts OfficeCLI. Resident capability projection +uses the native target canonically and advertises compatibility as a distinct +optional route. The preview MCP adapter has an explicit typed vocabulary rather than a command string passthrough. It supports validate, create/open/list, semantic get/query, @@ -417,7 +433,7 @@ component for one deprecation cycle before removal. Implemented: -1. Core, Browser, Office, extension, and component contracts. +1. Core, Browser, Office, OCR, extension, and component contracts. 2. Chrome and Lightpanda extraction from Search. 3. Search injection through `Arc`. 4. Typed Browser rendering and session tools over standard MCP stdio. @@ -479,6 +495,11 @@ Implemented: 15. A packaged first-party `a3s-use-office` Skill with progressive Word/Spreadsheet/Presentation/MCP references, bounded local discovery, release-archive smoke checks, and content-bound capability projection. +16. A first-party built-in OCR route with typed provider diagnostics, bounded + image admission, source SHA-256 evidence, local Tesseract and explicit + vision adapters, standard MCP annotations/output schemas, and a + release-packaged content-bound Skill that projects to `mcp__use_ocr__*` in + A3S Code. Next: diff --git a/src/capability_registry.rs b/src/capability_registry.rs index 92b57573..b8b02f27 100644 --- a/src/capability_registry.rs +++ b/src/capability_registry.rs @@ -76,6 +76,8 @@ pub(crate) async fn snapshot() -> UseResult { let mut capabilities = vec![ browser_capability().await?, office_capability().await?, + office_compatibility_capability(), + ocr_capability().await?, box_capability(), ]; capabilities.extend(extensions); @@ -167,26 +169,30 @@ async fn browser_capability() -> UseResult { async fn office_capability() -> UseResult { #[cfg(feature = "office")] { - let diagnostic = a3s_use_office::doctor(); - let ready = diagnostic.readiness == Readiness::Ready; let skill = crate::office_skills::primary_skill_surface().await; let (package_root, skills) = match skill { Some((root, path)) => (Some(root), vec![skill_surface(path).await?]), None => (None, Vec::new()), }; + let mut surfaces = vec!["cli".to_string(), "skill".to_string()]; + #[cfg(feature = "mcp")] + surfaces.push("mcp".to_string()); Ok(CapabilityBinding { id: "use/office".to_string(), route: "office".to_string(), version: env!("CARGO_PKG_VERSION").to_string(), origin: CapabilityOrigin::BuiltIn, enabled: true, - readiness: diagnostic.readiness, + readiness: Readiness::Ready, package_root, - surfaces: vec!["cli".to_string(), "mcp".to_string(), "skill".to_string()], - mcp: ready.then(|| McpSurface { - target: "office".to_string(), + surfaces, + #[cfg(feature = "mcp")] + mcp: Some(McpSurface { + target: "office-native".to_string(), transport: McpTransport::Stdio, }), + #[cfg(not(feature = "mcp"))] + mcp: None, skills, }) } @@ -207,6 +213,95 @@ async fn office_capability() -> UseResult { } } +fn office_compatibility_capability() -> CapabilityBinding { + #[cfg(feature = "office")] + { + let diagnostic = a3s_use_office::doctor(); + let ready = diagnostic.readiness == Readiness::Ready; + CapabilityBinding { + id: "use/office-compat".to_string(), + route: "office-compat".to_string(), + version: env!("CARGO_PKG_VERSION").to_string(), + origin: CapabilityOrigin::BuiltIn, + enabled: true, + readiness: diagnostic.readiness, + package_root: None, + surfaces: vec!["mcp".to_string()], + mcp: ready.then(|| McpSurface { + target: "office-compat".to_string(), + transport: McpTransport::Stdio, + }), + skills: Vec::new(), + } + } + #[cfg(not(feature = "office"))] + { + CapabilityBinding { + id: "use/office-compat".to_string(), + route: "office-compat".to_string(), + version: env!("CARGO_PKG_VERSION").to_string(), + origin: CapabilityOrigin::BuiltIn, + enabled: false, + readiness: Readiness::Missing, + package_root: None, + surfaces: Vec::new(), + mcp: None, + skills: Vec::new(), + } + } +} + +async fn ocr_capability() -> UseResult { + #[cfg(feature = "ocr")] + { + let diagnostic = crate::ocr_builtin::diagnostic(); + let skill = crate::ocr_builtin::primary_skill_surface().await; + let (package_root, skills) = match skill { + Some((root, path)) => (Some(root), vec![skill_surface(path).await?]), + None => (None, Vec::new()), + }; + let mut surfaces = vec!["cli".to_string()]; + if !skills.is_empty() { + surfaces.push("skill".to_string()); + } + #[cfg(feature = "mcp")] + surfaces.push("mcp".to_string()); + Ok(CapabilityBinding { + id: "use/ocr".to_string(), + route: "ocr".to_string(), + version: env!("CARGO_PKG_VERSION").to_string(), + origin: CapabilityOrigin::BuiltIn, + enabled: true, + readiness: diagnostic.readiness, + package_root, + surfaces, + #[cfg(feature = "mcp")] + mcp: Some(McpSurface { + target: "ocr-native".to_string(), + transport: McpTransport::Stdio, + }), + #[cfg(not(feature = "mcp"))] + mcp: None, + skills, + }) + } + #[cfg(not(feature = "ocr"))] + { + Ok(CapabilityBinding { + id: "use/ocr".to_string(), + route: "ocr".to_string(), + version: env!("CARGO_PKG_VERSION").to_string(), + origin: CapabilityOrigin::BuiltIn, + enabled: false, + readiness: Readiness::Missing, + package_root: None, + surfaces: Vec::new(), + mcp: None, + skills: Vec::new(), + }) + } +} + fn box_capability() -> CapabilityBinding { let diagnostic = crate::component_route::box_diagnostic(); CapabilityBinding { @@ -310,6 +405,13 @@ async fn project_extensions( ) -> UseResult>> { let mut capabilities = Vec::with_capacity(snapshot.routes.len()); for route in &snapshot.routes { + #[cfg(feature = "ocr")] + if route.route == "ocr" { + // OCR became a first-party built-in route. Ignore a legacy OCR + // extension receipt so an older installation cannot shadow or + // duplicate the release-matched built-in MCP/Skill projection. + continue; + } let Some(extension) = crate::extension_host::get(&route.package_id).await? else { return Ok(None); }; @@ -379,9 +481,21 @@ mod tests { .iter() .find(|capability| capability.id == "use/office") .unwrap(); + let office_compat = snapshot + .capabilities + .iter() + .find(|capability| capability.id == "use/office-compat") + .unwrap(); + let ocr = snapshot + .capabilities + .iter() + .find(|capability| capability.id == "use/ocr") + .unwrap(); assert_eq!(browser.origin, CapabilityOrigin::BuiltIn); assert_eq!(office.origin, CapabilityOrigin::BuiltIn); + assert_eq!(office_compat.origin, CapabilityOrigin::BuiltIn); + assert_eq!(ocr.origin, CapabilityOrigin::BuiltIn); #[cfg(feature = "browser")] { assert!(browser.surfaces.iter().any(|surface| surface == "skill")); @@ -405,6 +519,13 @@ mod tests { .iter() .any(|skill| skill.path.ends_with("a3s-use-office/SKILL.md"))); assert!(office.skills.iter().all(|skill| skill.sha256.len() == 64)); + #[cfg(feature = "mcp")] + assert_eq!( + office.mcp.as_ref().map(|surface| surface.target.as_str()), + Some("office-native") + ); + assert!(office_compat.skills.is_empty()); + assert_eq!(office_compat.route, "office-compat"); } #[cfg(not(feature = "office"))] { @@ -412,6 +533,27 @@ mod tests { assert!(office.surfaces.is_empty()); assert!(office.skills.is_empty()); } + #[cfg(feature = "ocr")] + { + assert!(ocr.enabled); + assert!(ocr.surfaces.iter().any(|surface| surface == "skill")); + assert!(ocr + .skills + .iter() + .any(|skill| skill.path.ends_with("a3s-use-ocr/SKILL.md"))); + assert!(ocr.skills.iter().all(|skill| skill.sha256.len() == 64)); + #[cfg(feature = "mcp")] + assert_eq!( + ocr.mcp.as_ref().map(|surface| surface.target.as_str()), + Some("ocr-native") + ); + } + #[cfg(not(feature = "ocr"))] + { + assert!(!ocr.enabled); + assert!(ocr.surfaces.is_empty()); + assert!(ocr.skills.is_empty()); + } assert_eq!(snapshot.revision.len(), 64); } diff --git a/src/cli.rs b/src/cli.rs index 07c6ad2a..23bad00f 100644 --- a/src/cli.rs +++ b/src/cli.rs @@ -63,6 +63,7 @@ pub async fn run(args: Vec) -> UseResult { "doctor" => doctor(args.get(1).map(String::as_str)), "component" => component(&args[1..]).await, "browser" => browser(&args[1..]).await, + "ocr" => ocr(&args[1..]).await, "box" => { let exit_code = crate::component_route::run_box(&args[1..]).await?; Ok(CommandOutput::delegated(exit_code)) @@ -90,6 +91,9 @@ fn version() -> CommandOutput { "schemaVersion": 1, "ok": true, "version": env!("CARGO_PKG_VERSION"), + "data": { + "version": env!("CARGO_PKG_VERSION"), + }, }), exit_code: 0, should_print: true, @@ -104,7 +108,7 @@ fn help() -> CommandOutput { " a3s-use capabilities [--json]\n", " a3s-use capability snapshot [--json]\n", " a3s-use capability watch [--after-generation ] [--after-revision ] [--timeout-ms ] [--json]\n", - " a3s-use doctor [browser|box|office] [--json]\n", + " a3s-use doctor [browser|box|office|ocr] [--json]\n", " a3s-use component list|status|install|uninstall [args] [--json]\n", " a3s-use browser doctor [--json]\n", " a3s-use browser render [--output ] [--screenshot ] [--json]\n", @@ -114,12 +118,14 @@ fn help() -> CommandOutput { " a3s-use office skills list|get|path [args] [--json]\n", " a3s-use office native get|query|view|watch|raw|raw-set|dump|merge|validate|create|add|add-part|set|sort|remove|move|copy|swap|insert-rows|delete-rows|insert-columns|delete-columns|rename-sheet|move-sheet|copy-sheet|batch [args] [--json]\n", " a3s-use office \n", + " a3s-use ocr doctor [--json]\n", + " a3s-use ocr extract [--language ] [--provider ] [--json]\n", " a3s-use extension list|inspect|doctor [args] [--json]\n", " a3s-use extension enable [--json]\n", " a3s-use extension disable [--timeout-ms ] [--json]\n", " a3s-use extension snapshot|watch [--after-generation ] [--timeout-ms ] [--json]\n", " a3s-use mcp serve browser [--tools ]\n", - " a3s-use mcp serve office|office-native|\n", + " a3s-use mcp serve office|office-native|office-compat|ocr|\n", " a3s-use mcp start|status|stop [browser] [--json]" ), serde_json::json!({ @@ -131,6 +137,7 @@ fn help() -> CommandOutput { "browser", "box", "office", + "ocr", "extension", "mcp" ] @@ -142,9 +149,10 @@ async fn capabilities() -> UseResult { let browser = browser_diagnostic(); let box_domain = crate::component_route::box_diagnostic(); let office = office_diagnostic(); + let ocr = ocr_diagnostic(); let (extension_generation, extensions) = extension_capabilities().await?; Ok(CommandOutput::success( - "Built-in routes: browser, box, office", + "Built-in routes: browser, box, office, ocr", serde_json::json!({ "domains": [ { @@ -159,6 +167,12 @@ async fn capabilities() -> UseResult { "readiness": office.readiness, "surfaces": ["cli", "mcp", "skill"] }, + { + "id": "ocr", + "builtIn": true, + "readiness": ocr.readiness, + "surfaces": ["cli", "mcp", "skill"] + }, { "id": "box", "builtIn": true, @@ -221,11 +235,13 @@ fn doctor(domain: Option<&str>) -> UseResult { None | Some("--json") => vec![ browser_diagnostic(), office_diagnostic(), + ocr_diagnostic(), crate::component_route::box_diagnostic(), ], Some("browser") => vec![browser_diagnostic()], Some("box") => vec![crate::component_route::box_diagnostic()], Some("office") => vec![office_diagnostic()], + Some("ocr") => vec![ocr_diagnostic()], Some(value) => { return Err(UseError::new( "use.domain_unknown", @@ -267,8 +283,9 @@ async fn component_list() -> UseResult { let browser = component_value("browser", &browser_diagnostic()); let box_component = component_value("box", &crate::component_route::box_diagnostic()); let office = component_value("office", &office_diagnostic()); + let ocr = component_value("ocr", &ocr_diagnostic()); let extensions = installed_extensions().await?; - let mut components = vec![browser, box_component, office]; + let mut components = vec![browser, box_component, office, ocr]; components.extend( extensions .iter() @@ -278,6 +295,7 @@ async fn component_list() -> UseResult { "browser".to_string(), "box".to_string(), "office".to_string(), + "ocr".to_string(), ]; human.extend( extensions @@ -495,7 +513,10 @@ async fn component_uninstall(id: &str) -> UseResult { )); } } - if matches!(id, "browser" | "use/browser" | "office" | "use/office") { + if matches!( + id, + "browser" | "use/browser" | "office" | "use/office" | "ocr" | "use/ocr" + ) { return Ok(CommandOutput::success( format!("No managed runtime files are owned for '{id}'."), serde_json::json!({ @@ -663,9 +684,11 @@ async fn mcp(args: &[String]) -> UseResult { "Standard Browser MCP support is disabled in this custom build.", )) } - "office" | "use/office" => { + "office" | "use/office" | "office-compat" | "use/office-compat" => { if args.len() != 2 { - return Err(usage_error("mcp serve office accepts exactly one target")); + return Err(usage_error( + "mcp serve office compatibility targets accept exactly one target", + )); } #[cfg(feature = "office")] { @@ -696,6 +719,21 @@ async fn mcp(args: &[String]) -> UseResult { "Native Office MCP support is disabled in this custom build.", )) } + "ocr" | "use/ocr" | "ocr-native" | "use/ocr-native" => { + if args.len() != 2 { + return Err(usage_error("mcp serve ocr accepts exactly one target")); + } + #[cfg(all(feature = "ocr", feature = "mcp"))] + { + a3s_use_ocr::OcrMcpServer::from_env()?.serve_stdio().await?; + Ok(CommandOutput::delegated(0)) + } + #[cfg(not(all(feature = "ocr", feature = "mcp")))] + Err(UseError::new( + "use.mcp.disabled", + "OCR MCP support is disabled in this custom build.", + )) + } package_id if external_package_id(package_id).is_some() => { if args.len() != 2 { return Err(usage_error( @@ -888,6 +926,7 @@ fn builtin_diagnostic(id: &str) -> Option { "browser" | "use/browser" => Some(browser_diagnostic()), "box" | "use/box" => Some(crate::component_route::box_diagnostic()), "office" | "use/office" => Some(office_diagnostic()), + "ocr" | "use/ocr" => Some(ocr_diagnostic()), _ => None, } } @@ -1031,7 +1070,21 @@ fn office_diagnostic() -> DomainDiagnostic { disabled_diagnostic("office") } -#[cfg(any(not(feature = "browser"), not(feature = "office")))] +#[cfg(feature = "ocr")] +fn ocr_diagnostic() -> DomainDiagnostic { + crate::ocr_builtin::diagnostic() +} + +#[cfg(not(feature = "ocr"))] +fn ocr_diagnostic() -> DomainDiagnostic { + disabled_diagnostic("ocr") +} + +#[cfg(any( + not(feature = "browser"), + not(feature = "office"), + not(feature = "ocr") +))] fn disabled_diagnostic(domain: &str) -> DomainDiagnostic { DomainDiagnostic { domain: domain.to_string(), @@ -1044,6 +1097,25 @@ fn disabled_diagnostic(domain: &str) -> DomainDiagnostic { } } +#[cfg(feature = "ocr")] +async fn ocr(args: &[String]) -> UseResult { + let output = a3s_use_ocr::cli::run(args.to_vec()).await?; + Ok(CommandOutput { + human: output.human, + json: output.json, + exit_code: output.exit_code, + should_print: output.should_print, + }) +} + +#[cfg(not(feature = "ocr"))] +async fn ocr(_args: &[String]) -> UseResult { + Err(UseError::new( + "use.ocr.disabled", + "OCR support is disabled in this custom build.", + )) +} + fn value_argument<'a>(args: &'a [String], index: usize, message: &str) -> UseResult<&'a str> { args.get(index) .map(String::as_str) diff --git a/src/cli_tests.rs b/src/cli_tests.rs index 6aa99587..e60610cd 100644 --- a/src/cli_tests.rs +++ b/src/cli_tests.rs @@ -1,13 +1,25 @@ use super::*; #[tokio::test] -async fn capabilities_always_include_browser_and_office() { +async fn version_json_exposes_a_typed_data_payload_for_consumers() { + let output = run(vec!["--version".to_string(), "--json".to_string()]) + .await + .unwrap(); + + assert_eq!(output.json["schemaVersion"], 1); + assert_eq!(output.json["ok"], true); + assert_eq!(output.json["data"]["version"], env!("CARGO_PKG_VERSION")); +} + +#[tokio::test] +async fn capabilities_always_include_browser_office_and_ocr() { let output = run(vec!["capabilities".to_string(), "--json".to_string()]) .await .unwrap(); let domains = output.json["data"]["domains"].as_array().unwrap(); assert_eq!(domains[0]["id"], "browser"); assert_eq!(domains[1]["id"], "office"); + assert_eq!(domains[2]["id"], "ocr"); assert!(domains[0]["surfaces"] .as_array() .unwrap() @@ -35,9 +47,14 @@ async fn capability_snapshot_unifies_built_ins_without_rpc_envelopes() { .iter() .find(|capability| capability["id"] == "use/office") .unwrap(); + let ocr = capabilities + .iter() + .find(|capability| capability["id"] == "use/ocr") + .unwrap(); assert_eq!(browser["origin"], "built-in"); assert_eq!(office["origin"], "built-in"); + assert_eq!(ocr["origin"], "built-in"); #[cfg(feature = "office")] { assert!(office["surfaces"] @@ -60,10 +77,41 @@ async fn capability_snapshot_unifies_built_ins_without_rpc_envelopes() { assert_eq!(office["surfaces"], serde_json::json!([])); assert!(office.get("skills").is_none()); } + #[cfg(feature = "ocr")] + { + assert_eq!(ocr["enabled"], true); + assert_eq!(ocr["mcp"]["target"], "ocr-native"); + assert!(ocr["skills"][0]["path"].as_str().is_some_and( + |path| std::path::Path::new(path).ends_with("skills/a3s-use-ocr/SKILL.md") + )); + assert_eq!(ocr["skills"][0]["sha256"].as_str().unwrap().len(), 64); + } + #[cfg(not(feature = "ocr"))] + { + assert_eq!(ocr["enabled"], false); + assert_eq!(ocr["surfaces"], serde_json::json!([])); + assert!(ocr.get("skills").is_none()); + } assert_eq!(registry["revision"].as_str().unwrap().len(), 64); assert!(output.json.get("jsonrpc").is_none()); } +#[cfg(feature = "ocr")] +#[tokio::test] +async fn built_in_ocr_doctor_uses_the_root_cli_contract() { + let output = run(vec![ + "ocr".to_string(), + "doctor".to_string(), + "--json".to_string(), + ]) + .await + .unwrap(); + + assert_eq!(output.json["schemaVersion"], 1); + assert_eq!(output.json["ok"], true); + assert!(output.json["data"]["readiness"].is_string()); +} + #[tokio::test] async fn component_status_uses_cli_json_contract() { let output = run(vec![ diff --git a/src/lib.rs b/src/lib.rs index 2c27ad51..fe1d9d97 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -11,6 +11,9 @@ pub mod cli; mod component_route; mod extension_cli; +#[cfg(feature = "ocr")] +mod ocr_builtin; + #[cfg(feature = "office")] mod office_artifact; #[cfg(feature = "office")] @@ -37,5 +40,8 @@ pub use a3s_use_browser as browser; #[cfg(feature = "office")] pub use a3s_use_office as office; +#[cfg(feature = "ocr")] +pub use a3s_use_ocr as ocr; + #[cfg(feature = "extensions")] pub use a3s_use_extension as extension; diff --git a/src/mcp.rs b/src/mcp.rs index 5c4c090c..68604713 100644 --- a/src/mcp.rs +++ b/src/mcp.rs @@ -210,7 +210,13 @@ mod browser { impl BrowserMcpServer { #[tool( name = "browser_doctor", - description = "Inspect the locally available A3S Use Browser provider without installing software" + description = "Inspect the locally available A3S Use Browser provider without installing software", + annotations( + read_only_hint = true, + destructive_hint = false, + idempotent_hint = true, + open_world_hint = false + ) )] async fn browser_doctor(&self) -> Result { Ok(match serde_json::to_value(a3s_use_browser::doctor()) { @@ -224,7 +230,13 @@ mod browser { #[tool( name = "browser_render", - description = "Render one web page with the configured local Browser provider" + description = "Render one web page with the configured local Browser provider", + annotations( + read_only_hint = true, + destructive_hint = false, + idempotent_hint = true, + open_world_hint = true + ) )] async fn browser_render( &self, @@ -249,7 +261,13 @@ mod browser { #[tool( name = "browser_open", - description = "Open an isolated stateful Browser session and return its first semantic snapshot" + description = "Open an isolated stateful Browser session and return its first semantic snapshot", + annotations( + read_only_hint = false, + destructive_hint = false, + idempotent_hint = false, + open_world_hint = true + ) )] async fn browser_open( &self, @@ -281,7 +299,13 @@ mod browser { #[tool( name = "browser_list", - description = "List open Browser sessions and their current URLs" + description = "List open Browser sessions and their current URLs", + annotations( + read_only_hint = true, + destructive_hint = false, + idempotent_hint = true, + open_world_hint = false + ) )] async fn browser_list(&self) -> Result { Ok(tool_result(self.sessions.list().await)) @@ -289,7 +313,13 @@ mod browser { #[tool( name = "browser_navigate", - description = "Navigate an open Browser session and return a fresh semantic snapshot" + description = "Navigate an open Browser session and return a fresh semantic snapshot", + annotations( + read_only_hint = false, + destructive_hint = false, + idempotent_hint = false, + open_world_hint = true + ) )] async fn browser_navigate( &self, @@ -320,7 +350,13 @@ mod browser { #[tool( name = "browser_snapshot", - description = "Return a compact semantic snapshot and fresh @e element references" + description = "Return a compact semantic snapshot and fresh @e element references", + annotations( + read_only_hint = true, + destructive_hint = false, + idempotent_hint = true, + open_world_hint = false + ) )] async fn browser_snapshot( &self, @@ -335,7 +371,13 @@ mod browser { #[tool( name = "browser_click", - description = "Click an element reference from the latest semantic snapshot" + description = "Click an element reference from the latest semantic snapshot", + annotations( + read_only_hint = false, + destructive_hint = false, + idempotent_hint = false, + open_world_hint = true + ) )] async fn browser_click( &self, @@ -352,7 +394,13 @@ mod browser { #[tool( name = "browser_type", - description = "Focus an element reference and type text into it" + description = "Focus an element reference and type text into it", + annotations( + read_only_hint = false, + destructive_hint = false, + idempotent_hint = false, + open_world_hint = true + ) )] async fn browser_type( &self, @@ -371,7 +419,13 @@ mod browser { #[tool( name = "browser_press", - description = "Focus an element reference and press one keyboard key" + description = "Focus an element reference and press one keyboard key", + annotations( + read_only_hint = false, + destructive_hint = false, + idempotent_hint = false, + open_world_hint = true + ) )] async fn browser_press( &self, @@ -390,7 +444,13 @@ mod browser { #[tool( name = "browser_select", - description = "Select an option value on a referenced select element" + description = "Select an option value on a referenced select element", + annotations( + read_only_hint = false, + destructive_hint = false, + idempotent_hint = false, + open_world_hint = true + ) )] async fn browser_select( &self, @@ -409,7 +469,13 @@ mod browser { #[tool( name = "browser_scroll", - description = "Scroll the current page by explicit horizontal and vertical deltas" + description = "Scroll the current page by explicit horizontal and vertical deltas", + annotations( + read_only_hint = false, + destructive_hint = false, + idempotent_hint = false, + open_world_hint = false + ) )] async fn browser_scroll( &self, @@ -426,7 +492,13 @@ mod browser { #[tool( name = "browser_screenshot", - description = "Capture a full-page PNG from an open Browser session to an explicit local path" + description = "Capture a full-page PNG from an open Browser session to an explicit local path", + annotations( + read_only_hint = false, + destructive_hint = false, + idempotent_hint = false, + open_world_hint = false + ) )] async fn browser_screenshot( &self, @@ -449,7 +521,13 @@ mod browser { #[tool( name = "browser_close", - description = "Close one Browser session and release its tab resources" + description = "Close one Browser session and release its tab resources", + annotations( + read_only_hint = false, + destructive_hint = false, + idempotent_hint = true, + open_world_hint = false + ) )] async fn browser_close( &self, @@ -470,7 +548,13 @@ mod browser { #[tool( name = "browser_service_stop", - description = "Stop the authenticated persistent A3S Use Browser MCP deployment" + description = "Stop the authenticated persistent A3S Use Browser MCP deployment", + annotations( + read_only_hint = false, + destructive_hint = true, + idempotent_hint = true, + open_world_hint = false + ) )] async fn browser_service_stop(&self) -> Result { let Some(shutdown) = self.shutdown.clone() else { @@ -564,7 +648,7 @@ mod browser { None, ); let tools = server.tool_router.list_all(); - let mut names = tools + let mut names: Vec<&str> = tools .iter() .map(|tool| tool.name.as_ref()) .collect::>(); @@ -587,6 +671,23 @@ mod browser { "browser_type" ] ); + + let annotations = |name: &str| { + tools + .iter() + .find(|tool| tool.name == name) + .and_then(|tool| tool.annotations.as_ref()) + .unwrap_or_else(|| panic!("{name} must declare MCP annotations")) + }; + let list = annotations("browser_list"); + assert_eq!(list.read_only_hint, Some(true)); + assert_eq!(list.open_world_hint, Some(false)); + let render = annotations("browser_render"); + assert_eq!(render.read_only_hint, Some(true)); + assert_eq!(render.open_world_hint, Some(true)); + let click = annotations("browser_click"); + assert_eq!(click.read_only_hint, Some(false)); + assert_eq!(click.open_world_hint, Some(true)); } #[test] diff --git a/src/mcp/office.rs b/src/mcp/office.rs index 1b84ff54..3305ff59 100644 --- a/src/mcp/office.rs +++ b/src/mcp/office.rs @@ -86,7 +86,13 @@ impl NativeOfficeMcpServer { impl NativeOfficeMcpServer { #[tool( name = "office_validate", - description = "Validate and identify one local OOXML document without opening a session" + description = "Validate and identify one local OOXML document without opening a session", + annotations( + read_only_hint = true, + destructive_hint = false, + idempotent_hint = true, + open_world_hint = false + ) )] async fn office_validate( &self, @@ -108,7 +114,13 @@ impl NativeOfficeMcpServer { #[tool( name = "office_create", - description = "Create a blank native OOXML document and register a mutable in-memory session" + description = "Create a blank native OOXML document and register a mutable in-memory session", + annotations( + read_only_hint = false, + destructive_hint = false, + idempotent_hint = false, + open_world_hint = false + ) )] async fn office_create( &self, @@ -125,7 +137,13 @@ impl NativeOfficeMcpServer { #[tool( name = "office_open", - description = "Open a local OOXML document in a bounded native in-memory session" + description = "Open a local OOXML document in a bounded native in-memory session", + annotations( + read_only_hint = false, + destructive_hint = false, + idempotent_hint = false, + open_world_hint = false + ) )] async fn office_open( &self, @@ -145,7 +163,13 @@ impl NativeOfficeMcpServer { #[tool( name = "office_list", - description = "List native Office sessions owned by this MCP server process" + description = "List native Office sessions owned by this MCP server process", + annotations( + read_only_hint = true, + destructive_hint = false, + idempotent_hint = true, + open_world_hint = false + ) )] async fn office_list(&self) -> Result { let mut entries = self.sessions.list().await; @@ -165,7 +189,13 @@ impl NativeOfficeMcpServer { #[tool( name = "office_get", - description = "Read one stable semantic path from an open native Office session" + description = "Read one stable semantic path from an open native Office session", + annotations( + read_only_hint = true, + destructive_hint = false, + idempotent_hint = true, + open_world_hint = false + ) )] async fn office_get( &self, @@ -192,7 +222,13 @@ impl NativeOfficeMcpServer { #[tool( name = "office_query", - description = "Run a native semantic selector with a bounded result count" + description = "Run a native semantic selector with a bounded result count", + annotations( + read_only_hint = true, + destructive_hint = false, + idempotent_hint = true, + open_world_hint = false + ) )] async fn office_query( &self, @@ -228,7 +264,13 @@ impl NativeOfficeMcpServer { #[tool( name = "office_view", - description = "Produce a native text, bounded annotated, outline, statistics, bounded issues, standalone all-format HTML or SVG, or Browser-injected PNG screenshot view for an open session" + description = "Produce a native text, bounded annotated, outline, statistics, bounded issues, standalone all-format HTML or SVG, or Browser-injected PNG screenshot view for an open session", + annotations( + read_only_hint = false, + destructive_hint = false, + idempotent_hint = false, + open_world_hint = false + ) )] async fn office_view( &self, @@ -303,7 +345,13 @@ impl NativeOfficeMcpServer { #[tool( name = "office_raw_xml", - description = "Inspect one existing OOXML XML part, limited to 1 MiB of original bytes" + description = "Inspect one existing OOXML XML part, limited to 1 MiB of original bytes", + annotations( + read_only_hint = true, + destructive_hint = false, + idempotent_hint = true, + open_world_hint = false + ) )] async fn office_raw_xml( &self, @@ -332,7 +380,13 @@ impl NativeOfficeMcpServer { #[tool( name = "office_apply_batch", - description = "Apply a bounded typed mutation batch atomically in memory; call office_save to persist it" + description = "Apply a bounded typed mutation batch atomically in memory; call office_save to persist it", + annotations( + read_only_hint = false, + destructive_hint = false, + idempotent_hint = false, + open_world_hint = false + ) )] async fn office_apply_batch( &self, @@ -363,7 +417,13 @@ impl NativeOfficeMcpServer { #[tool( name = "office_merge_template", - description = "Merge bounded JSON data into a cloned session document and atomically save a distinct output" + description = "Merge bounded JSON data into a cloned session document and atomically save a distinct output", + annotations( + read_only_hint = false, + destructive_hint = false, + idempotent_hint = false, + open_world_hint = false + ) )] async fn office_merge_template( &self, @@ -403,7 +463,13 @@ impl NativeOfficeMcpServer { #[tool( name = "office_save", - description = "Atomically persist one mutable native Office session, optionally to a new path" + description = "Atomically persist one mutable native Office session, optionally to a new path", + annotations( + read_only_hint = false, + destructive_hint = true, + idempotent_hint = false, + open_world_hint = false + ) )] async fn office_save( &self, @@ -430,7 +496,13 @@ impl NativeOfficeMcpServer { #[tool( name = "office_close", - description = "Close a native Office session, refusing unsaved changes unless discard is explicit" + description = "Close a native Office session, refusing unsaved changes unless discard is explicit", + annotations( + read_only_hint = false, + destructive_hint = true, + idempotent_hint = true, + open_world_hint = false + ) )] async fn office_close( &self, diff --git a/src/mcp/office/tests.rs b/src/mcp/office/tests.rs index 10827174..8e70566a 100644 --- a/src/mcp/office/tests.rs +++ b/src/mcp/office/tests.rs @@ -14,7 +14,7 @@ use a3s_use_office::{ fn native_office_server_exposes_only_bounded_typed_tools() { let server = NativeOfficeMcpServer::new(); let tools = server.tool_router.list_all(); - let mut names = tools + let mut names: Vec<&str> = tools .iter() .map(|tool| tool.name.as_ref()) .collect::>(); @@ -36,6 +36,30 @@ fn native_office_server_exposes_only_bounded_typed_tools() { "office_view", ] ); + + let annotations = |name: &str| { + tools + .iter() + .find(|tool| tool.name == name) + .and_then(|tool| tool.annotations.as_ref()) + .unwrap_or_else(|| panic!("{name} must declare MCP annotations")) + }; + for name in [ + "office_validate", + "office_list", + "office_get", + "office_query", + "office_raw_xml", + ] { + let annotation = annotations(name); + assert_eq!(annotation.read_only_hint, Some(true), "{name}"); + assert_eq!(annotation.open_world_hint, Some(false), "{name}"); + } + assert_eq!( + annotations("office_apply_batch").read_only_hint, + Some(false) + ); + assert_eq!(annotations("office_save").destructive_hint, Some(true)); } #[test] diff --git a/src/ocr_builtin.rs b/src/ocr_builtin.rs new file mode 100644 index 00000000..127004c2 --- /dev/null +++ b/src/ocr_builtin.rs @@ -0,0 +1,95 @@ +//! Built-in projection glue for the first-party OCR domain. + +use std::path::{Path, PathBuf}; + +use a3s_use_core::{DomainDiagnostic, Readiness}; +use a3s_use_ocr::{OcrClient, OcrProviderKind}; + +pub(crate) fn diagnostic() -> DomainDiagnostic { + match OcrClient::from_env() { + Ok(client) => { + let diagnostic = client.diagnostic(); + DomainDiagnostic { + domain: "ocr".to_string(), + readiness: diagnostic.readiness, + provider: diagnostic.provider.map(provider_name).map(str::to_string), + version: None, + path: diagnostic.executable, + message: diagnostic.message, + suggestions: diagnostic.suggestions, + } + } + Err(error) => DomainDiagnostic { + domain: "ocr".to_string(), + readiness: Readiness::Broken, + provider: None, + version: None, + path: None, + message: error.message, + suggestions: error.suggestion.into_iter().collect(), + }, + } +} + +pub(crate) async fn primary_skill_surface() -> Option<(PathBuf, PathBuf)> { + let mut roots = Vec::new(); + if let Some(root) = std::env::var_os("A3S_USE_OCR_SKILLS_DIR").map(PathBuf::from) { + roots.push(root); + } + if let Ok(executable) = std::env::current_exe() { + if let Some(parent) = executable.parent() { + roots.push(parent.join("ocr-skills")); + } + } + roots.push( + Path::new(env!("CARGO_MANIFEST_DIR")) + .join("crates") + .join("ocr") + .join("skills"), + ); + + for root in roots { + let skill = root.join("a3s-use-ocr/SKILL.md"); + let Ok(root) = tokio::fs::canonicalize(root).await else { + continue; + }; + let Ok(skill) = tokio::fs::canonicalize(skill).await else { + continue; + }; + if skill.starts_with(&root) { + return Some((root, skill)); + } + } + None +} + +fn provider_name(provider: OcrProviderKind) -> &'static str { + match provider { + OcrProviderKind::Auto => "auto", + OcrProviderKind::Tesseract => "tesseract", + OcrProviderKind::Vision => "vision", + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[tokio::test] + async fn source_skill_is_available_to_development_builds() { + let (root, skill) = primary_skill_surface().await.unwrap(); + assert!(root.is_absolute()); + assert!(skill.starts_with(root)); + assert!(skill.ends_with("a3s-use-ocr/SKILL.md")); + } + + #[test] + fn diagnostic_is_typed_even_without_a_provider() { + let diagnostic = diagnostic(); + assert_eq!(diagnostic.domain, "ocr"); + assert!(matches!( + diagnostic.readiness, + Readiness::Ready | Readiness::Missing | Readiness::Broken + )); + } +} diff --git a/tests/cli.rs b/tests/cli.rs index 3bdb785f..f502594b 100644 --- a/tests/cli.rs +++ b/tests/cli.rs @@ -21,9 +21,13 @@ fn capabilities_are_available_as_versioned_json() { assert!(output.status.success()); let value: serde_json::Value = serde_json::from_slice(&output.stdout).unwrap(); assert_eq!(value["schemaVersion"], 1); - assert_eq!(value["data"]["domains"][0]["id"], "browser"); - assert_eq!(value["data"]["domains"][1]["id"], "office"); - assert_eq!(value["data"]["domains"][2]["id"], "box"); + let domains = value["data"]["domains"].as_array().unwrap(); + for id in ["browser", "office", "ocr", "box"] { + assert!( + domains.iter().any(|domain| domain["id"] == id), + "missing built-in domain {id}: {domains:?}" + ); + } assert!(value["data"].get("customJsonRpc").is_none()); assert!(value.get("jsonrpc").is_none()); } @@ -58,6 +62,10 @@ fn unified_capability_snapshot_projects_builtin_skills() { .iter() .find(|capability| capability["id"] == "use/office") .unwrap(); + let office_compat = capabilities + .iter() + .find(|capability| capability["id"] == "use/office-compat") + .unwrap(); assert_eq!(browser["origin"], "built-in"); #[cfg(feature = "browser")] { @@ -100,6 +108,15 @@ fn unified_capability_snapshot_projects_builtin_skills() { assert!(office_skill_digest .bytes() .all(|byte| byte.is_ascii_hexdigit() && !byte.is_ascii_uppercase())); + #[cfg(feature = "mcp")] + { + assert_eq!(office["readiness"], "ready"); + assert_eq!(office["mcp"]["target"], "office-native"); + } + assert_eq!(office_compat["route"], "office-compat"); + assert_eq!(office_compat["readiness"], "missing"); + assert!(office_compat.get("mcp").is_none()); + assert!(office_compat.get("skills").is_none()); } #[cfg(not(feature = "office"))] { @@ -152,6 +169,9 @@ fn office_skill_commands_are_packaged_and_provider_independent() { .as_str() .unwrap() .contains("## Bundled reference: references/mcp.md")); + let office_skill = get["data"]["content"].as_str().unwrap(); + assert!(office_skill.contains("mcp__use_office__*")); + assert!(office_skill.contains("mcp__use_office_compat__*")); let path = Command::new(binary()) .args(["office", "skills", "path", "a3s-use-office", "--json"]) @@ -1811,6 +1831,27 @@ fn office_mcp_target_delegates_to_officeclis_standard_server() { assert!(output.stderr.is_empty()); } +#[cfg(all(unix, feature = "office"))] +#[test] +fn office_compat_mcp_target_delegates_to_officeclis_standard_server() { + let temp = tempfile::tempdir().unwrap(); + let executable = temp.path().join("officecli-fixture"); + std::fs::write(&executable, "#!/bin/sh\nprintf '%s\\n' \"$*\"\nexit 5\n").unwrap(); + let mut permissions = std::fs::metadata(&executable).unwrap().permissions(); + permissions.set_mode(0o755); + std::fs::set_permissions(&executable, permissions).unwrap(); + + let output = Command::new(binary()) + .args(["mcp", "serve", "office-compat"]) + .env("A3S_OFFICECLI_EXECUTABLE", &executable) + .output() + .unwrap(); + + assert_eq!(output.status.code(), Some(5)); + assert_eq!(String::from_utf8(output.stdout).unwrap(), "mcp\n"); + assert!(output.stderr.is_empty()); +} + #[cfg(all(feature = "office", feature = "mcp"))] #[tokio::test] async fn native_office_mcp_is_standard_typed_and_independent_of_officecli() { @@ -2100,6 +2141,36 @@ async fn standard_mcp_request( serde_json::from_str(&line).unwrap() } +#[cfg(feature = "ocr")] +#[test] +fn built_in_ocr_projects_the_canonical_code_route_and_skill() { + let temp = tempfile::tempdir().unwrap(); + let home = temp.path().join("home"); + + let snapshot = Command::new(binary()) + .args(["capability", "snapshot", "--json"]) + .env("A3S_USE_HOME", &home) + .output() + .unwrap(); + assert!(snapshot.status.success(), "{snapshot:?}"); + let snapshot: serde_json::Value = serde_json::from_slice(&snapshot.stdout).unwrap(); + let ocr = snapshot["data"]["registry"]["capabilities"] + .as_array() + .unwrap() + .iter() + .find(|capability| capability["id"] == "use/ocr") + .unwrap(); + assert_eq!(ocr["route"], "ocr"); + assert_eq!(ocr["origin"], "built-in"); + assert_eq!(ocr["enabled"], true); + assert_eq!(ocr["mcp"]["target"], "ocr-native"); + assert!(ocr["skills"][0]["path"] + .as_str() + .is_some_and(|path| Path::new(path).ends_with("skills/a3s-use-ocr/SKILL.md"))); + let digest = ocr["skills"][0]["sha256"].as_str().unwrap(); + assert_eq!(digest.len(), 64); +} + #[cfg(all(unix, feature = "extensions"))] #[test] fn explicit_extension_install_delegates_native_cli_and_preserves_status() { From 1b2a60d9c8930dcc9935b6841ed4a7778111f884 Mon Sep 17 00:00:00 2001 From: RoyLin Date: Sun, 19 Jul 2026 08:53:57 +0800 Subject: [PATCH 2/9] feat(extension): add signed remote registries --- Cargo.lock | 258 +++++- Cargo.toml | 7 + README.md | 74 +- crates/extension/Cargo.toml | 10 + crates/extension/src/digest.rs | 194 +++++ crates/extension/src/lib.rs | 7 + crates/extension/src/package.rs | 6 +- crates/extension/src/paths.rs | 10 + crates/extension/src/registry.rs | 251 +++++- crates/extension/src/registry_tests.rs | 252 +++++- crates/extension/src/remote.rs | 970 +++++++++++++++++++++++ crates/extension/src/remote_tests.rs | 310 ++++++++ crates/extension/src/source.rs | 708 +++++++++++++++++ crates/extension/src/tuf_test_support.rs | 324 ++++++++ docs/architecture.md | 18 +- src/cli.rs | 95 ++- src/extension_cli.rs | 78 +- src/extension_host.rs | 22 +- tests/extension_archives.rs | 102 +++ tests/remote_extension_cli.rs | 183 +++++ 20 files changed, 3829 insertions(+), 50 deletions(-) create mode 100644 crates/extension/src/digest.rs create mode 100644 crates/extension/src/remote.rs create mode 100644 crates/extension/src/remote_tests.rs create mode 100644 crates/extension/src/source.rs create mode 100644 crates/extension/src/tuf_test_support.rs create mode 100644 tests/extension_archives.rs create mode 100644 tests/remote_extension_cli.rs diff --git a/Cargo.lock b/Cargo.lock index 80e0f846..07033fd1 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -21,15 +21,19 @@ dependencies = [ "axum", "base64", "clap", + "flate2", "fs2", "futures-util", "getrandom 0.3.4", + "olpc-cjson", "reqwest", + "ring", "rmcp", "schemars", "serde", "serde_json", "sha2 0.10.9", + "tar", "tempfile", "tokio", "tokio-util", @@ -107,13 +111,21 @@ version = "0.1.1" dependencies = [ "a3s-acl", "a3s-use-core", + "flate2", "fs2", + "olpc-cjson", + "reqwest", + "ring", "semver", "serde", "serde_json", "sha2 0.10.9", + "tar", "tempfile", "tokio", + "tough", + "url", + "zip", ] [[package]] @@ -432,6 +444,17 @@ dependencies = [ "rustix 1.1.4", ] +[[package]] +name = "async-recursion" +version = "1.1.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "3b43422f69d8ff38f95f1b2bb76517c91589a924d1559a0e935d7c8ce0274c11" +dependencies = [ + "proc-macro2", + "quote", + "syn 2.0.118", +] + [[package]] name = "async-signal" version = "0.2.14" @@ -565,6 +588,30 @@ dependencies = [ "arrayvec", ] +[[package]] +name = "aws-lc-rs" +version = "1.17.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "00bdb5da18dac48ca2cc7cd4a98e533e8635a58e2361d13a1a4ee3888e0d72f1" +dependencies = [ + "aws-lc-sys", + "untrusted 0.7.1", + "zeroize", +] + +[[package]] +name = "aws-lc-sys" +version = "0.43.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "43103168cc76fe62678a375e722fc9cb3a0146159ac5828bc4f0dfd755c2224c" +dependencies = [ + "cc", + "cmake", + "dunce", + "fs_extra", + "pkg-config", +] + [[package]] name = "axum" version = "0.8.9" @@ -675,6 +722,16 @@ dependencies = [ "piper", ] +[[package]] +name = "bstr" +version = "1.13.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1f7dc094d718f2e1c1559ad110e27eeaae14a5465d3d56dd6dbd793079fbd530" +dependencies = [ + "memchr", + "serde_core", +] + [[package]] name = "built" version = "0.8.1" @@ -881,6 +938,15 @@ version = "1.1.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "c8d4a3bb8b1e0c1050499d1815f5ab16d04f0959b233085fb31653fbfc9d98f9" +[[package]] +name = "cmake" +version = "0.1.58" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c0f78a02292a74a88ac736019ab962ece0bc380e3f977bf72e376c5d78ff0678" +dependencies = [ + "cc", +] + [[package]] name = "color_quant" version = "1.1.0" @@ -1238,6 +1304,16 @@ dependencies = [ "simd-adler32", ] +[[package]] +name = "filetime" +version = "0.2.29" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "5c287a33c7f0a620c38e641e7f60827713987b3c0f26e8ddc9462cc69cf75759" +dependencies = [ + "cfg-if", + "libc", +] + [[package]] name = "find-msvc-tools" version = "0.1.9" @@ -1279,6 +1355,12 @@ dependencies = [ "winapi", ] +[[package]] +name = "fs_extra" +version = "1.3.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "42703706b716c37f96a77aea830392ad231f44c9e9a67872fa5548707e11b11c" + [[package]] name = "futures" version = "0.3.32" @@ -1455,6 +1537,19 @@ dependencies = [ "weezl", ] +[[package]] +name = "globset" +version = "0.4.19" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e47d37d2ae4464254884b60ab7071be2b876a9c35b696bd018ddcc76847309cd" +dependencies = [ + "aho-corasick", + "bstr", + "log", + "regex-automata", + "regex-syntax", +] + [[package]] name = "gloo-timers" version = "0.3.0" @@ -2144,6 +2239,17 @@ dependencies = [ "autocfg", ] +[[package]] +name = "olpc-cjson" +version = "0.1.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "696183c9b5fe81a7715d074fd632e8bd46f4ccc0231a3ed7fc580a80de5f7083" +dependencies = [ + "serde", + "serde_json", + "unicode-normalization", +] + [[package]] name = "once_cell" version = "1.21.4" @@ -2186,12 +2292,42 @@ version = "0.1.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "35fb2e5f958ec131621fdd531e9fc186ed768cbe395337403ae56c17a74c68ec" +[[package]] +name = "pem" +version = "3.0.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1d30c53c26bc5b31a98cd02d20f25a7c8567146caf63ed593a9d87b2775291be" +dependencies = [ + "base64", + "serde_core", +] + [[package]] name = "percent-encoding" version = "2.3.2" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "9b4f627cb1b25917193a259e49bdad08f671f8d9708acfd5fe0a8c1455d87220" +[[package]] +name = "pin-project" +version = "1.1.13" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "2466b2336ed02bcdca6b294417127b90ec92038d1d5c4fbeac971a922e0e0924" +dependencies = [ + "pin-project-internal", +] + +[[package]] +name = "pin-project-internal" +version = "1.1.13" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c96395f0a926bc13b1c17622aaddda1ecb55d49c8f1bf9777e4d877800a43f8b" +dependencies = [ + "proc-macro2", + "quote", + "syn 2.0.118", +] + [[package]] name = "pin-project-lite" version = "0.2.17" @@ -2215,6 +2351,12 @@ dependencies = [ "futures-io", ] +[[package]] +name = "pkg-config" +version = "0.3.33" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "19f132c84eca552bf34cab8ec81f1c1dcc229b811638f9d283dceabe58c5569e" + [[package]] name = "png" version = "0.18.1" @@ -2723,7 +2865,7 @@ dependencies = [ "cfg-if", "getrandom 0.2.17", "libc", - "untrusted", + "untrusted 0.9.0", "windows-sys 0.52.0", ] @@ -2844,6 +2986,8 @@ version = "0.23.42" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "3c54fcab019b409d04215d3a17cb438fd7fbf192ee61461f20f4fe18704bc138" dependencies = [ + "aws-lc-rs", + "log", "once_cell", "ring", "rustls-pki-types", @@ -2868,9 +3012,10 @@ version = "0.103.13" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "61c429a8649f110dddef65e2a5ad240f747e85f7758a6bccc7e5777bd33f756e" dependencies = [ + "aws-lc-rs", "ring", "rustls-pki-types", - "untrusted", + "untrusted 0.9.0", ] [[package]] @@ -2991,6 +3136,15 @@ dependencies = [ "serde_core", ] +[[package]] +name = "serde_plain" +version = "1.0.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9ce1fc6db65a611022b23a0dec6975d63fb80a302cb3388835ff02c097258d50" +dependencies = [ + "serde", +] + [[package]] name = "serde_urlencoded" version = "0.7.1" @@ -3085,6 +3239,29 @@ version = "1.15.2" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "8ed6a63f02c8539c91a8685a86f4099661ba3da017932f6ebbea6de3f0fa7c90" +[[package]] +name = "snafu" +version = "0.8.9" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "6e84b3f4eacbf3a1ce05eac6763b4d629d60cbc94d632e4092c54ade71f1e1a2" +dependencies = [ + "futures-core", + "pin-project", + "snafu-derive", +] + +[[package]] +name = "snafu-derive" +version = "0.8.9" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c1c97747dbf44bb1ca44a561ece23508e99cb592e862f22222dcf42f51d1e451" +dependencies = [ + "heck 0.5.0", + "proc-macro2", + "quote", + "syn 2.0.118", +] + [[package]] name = "socket2" version = "0.6.5" @@ -3168,6 +3345,17 @@ dependencies = [ "syn 2.0.118", ] +[[package]] +name = "tar" +version = "0.4.46" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "3f6221d9a6003c78398e3b239969f352578258df48c8eb051caadae0015bc840" +dependencies = [ + "filetime", + "libc", + "xattr", +] + [[package]] name = "tempfile" version = "3.27.0" @@ -3367,6 +3555,41 @@ dependencies = [ "tokio", ] +[[package]] +name = "tough" +version = "0.22.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8031cff0872dd1c6312370515a6be8098f6ea5512f1bad725016046fc725f272" +dependencies = [ + "async-recursion", + "async-trait", + "aws-lc-rs", + "bytes", + "chrono", + "dyn-clone", + "futures", + "futures-core", + "globset", + "hex", + "log", + "olpc-cjson", + "pem", + "percent-encoding", + "reqwest", + "rustls", + "serde", + "serde_json", + "serde_plain", + "snafu", + "tempfile", + "tokio", + "tokio-util", + "typed-path", + "untrusted 0.7.1", + "url", + "walkdir", +] + [[package]] name = "tower" version = "0.5.3" @@ -3489,6 +3712,12 @@ dependencies = [ "utf-8", ] +[[package]] +name = "typed-path" +version = "0.9.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "82205ffd44a9697e34fc145491aa47310f9871540bb7909eaa9365e0a9a46607" + [[package]] name = "typenum" version = "1.20.1" @@ -3507,6 +3736,15 @@ version = "1.0.24" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "e6e4313cd5fcd3dad5cafa179702e2b244f760991f45397d14d4ebf38247da75" +[[package]] +name = "unicode-normalization" +version = "0.1.25" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "5fd4f6878c9cb28d874b009da9e8d183b5abc80117c40bbd187a1fde336be6e8" +dependencies = [ + "tinyvec", +] + [[package]] name = "universal-hash" version = "0.5.1" @@ -3517,6 +3755,12 @@ dependencies = [ "subtle", ] +[[package]] +name = "untrusted" +version = "0.7.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a156c684c91ea7d62626509bce3cb4e1d9ed5c4d978f7b4352658f96a4c26b4a" + [[package]] name = "untrusted" version = "0.9.0" @@ -4039,6 +4283,16 @@ version = "0.6.3" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "1ffae5123b2d3fc086436f8834ae3ab053a283cfac8fe0a0b8eaae044768a4c4" +[[package]] +name = "xattr" +version = "1.6.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "32e45ad4206f6d2479085147f02bc2ef834ac85886624a23575ae137c8aa8156" +dependencies = [ + "libc", + "rustix 1.1.4", +] + [[package]] name = "y4m" version = "0.8.0" diff --git a/Cargo.toml b/Cargo.toml index 8c17bb9f..cb782059 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -25,6 +25,7 @@ base64 = "0.22" clap = { version = "4", features = ["derive"] } fs2 = "0.4" futures-util = "0.3" +flate2 = "1" getrandom = "0.3" reqwest = { version = "0.12", default-features = false, features = ["rustls-tls", "stream"] } quick-xml = "0.38" @@ -34,10 +35,12 @@ schemars = "1.2" serde = { version = "1", features = ["derive"] } serde_json = "1" sha2 = "0.10" +tar = "0.4" thiserror = "2" tempfile = "3" tokio = { version = "1", features = ["fs", "io-util", "macros", "net", "rt-multi-thread", "process", "sync", "time"] } tokio-util = "0.7" +tough = { version = "0.22", default-features = false, features = ["http"] } url = "2" zip = { version = "2", default-features = false, features = ["deflate"] } @@ -111,5 +114,9 @@ windows-sys = { version = "0.52", features = ["Win32_Foundation", "Win32_System_ [dev-dependencies] async-trait.workspace = true reqwest.workspace = true +flate2.workspace = true +olpc-cjson = "0.1" +ring = "0.17" +tar.workspace = true tempfile.workspace = true zip.workspace = true diff --git a/README.md b/README.md index 5b4916fa..9304819c 100644 --- a/README.md +++ b/README.md @@ -1631,6 +1631,7 @@ the parent TUI before source bytes leave the device. See the [OCR crate](crates/ocr/README.md) for configuration and provider boundaries. + ## External Extensions External Use domains stay behind process boundaries. A package contains an @@ -1674,11 +1675,74 @@ a3s use extension enable acme/slack --json a3s uninstall use/acme/slack ``` -The current extension source is an explicit local directory. It must pass -manifest, route, path, package-size, and executable validation, and unsigned -content requires `--allow-unsigned`. A signed remote publisher channel is -roadmap work; Use does not silently install arbitrary Homebrew, npm, Cargo, -system, or `PATH` packages. +The current extension source is an explicit local directory or a `.tar.gz`, +`.tgz`, or `.zip` archive. Archives must contain exactly one package manifest; +every entry must belong to that manifest's package root. Installation rejects +links, traversal, duplicate paths, unsupported entries, excessive expansion, +and non-portable paths before validating the manifest, route, executable, and +Skill surfaces. Unsigned content requires `--allow-unsigned`. Use does not +silently install arbitrary Homebrew, npm, Cargo, system, or `PATH` packages. + +### Signed extension registries + +Remote extensions use TUF metadata and a separately established bootstrap-root +digest. Enroll a registry with either a root file or its SHA-256, verify it, +review the immutable component plan, and apply that exact plan: + +```bash +a3s registry add https://packages.example.org/a3s/ \ + --trust-root ./root.json \ + --yes +a3s registry refresh packages + +a3s --output json install use/acme/slack --dry-run +a3s --output json install use/acme/slack \ + --plan-digest + +a3s --output json upgrade use/acme/slack --dry-run +a3s --output json upgrade use/acme/slack \ + --plan-digest +``` + +When a root file is supplied, the umbrella CLI copies it into registry-owned +configuration and records its digest. With a digest-only enrollment, Use may +fetch `/metadata/root.json`, but it caches the file only after the +bytes match the pinned SHA-256. Subsequent root rotation, timestamp, snapshot, +and targets metadata are verified by TUF with expiration and rollback +enforcement. Registry URLs require HTTPS; loopback HTTP is accepted only for +tests and local development. + +A dry-run verifies metadata but does not download the target archive. Its outer +component digest includes the exact `ResolvedRemotePackage`: registry identity, +bootstrap root, every TUF metadata version, package version and channel, +platform target, archive path, length, and SHA-256. Apply resolves again and +fails before target download if that plan changed. It then passes the resolved +package's own digest to `a3s-use`, which repeats TUF verification immediately +before downloading and activating the archive. The installed receipt records +`registry-tuf` trust and the complete signed provenance. Registry installs +reject `--allow-unsigned`; local `--from` installs cannot provide registry +options. + +Registry upgrades reuse the registry identity and channel recorded in that +signed provenance instead of searching every configured source again. A +missing registry, changed URL or bootstrap root, and semantic-version downgrade +are rejected before payload download. Plain `a3s upgrade` reports newer signed +targets, while `a3s upgrade --all` includes them in the selected batch. If the +verified target is identical to the installed target, `a3s-use` validates and +reconciles the receipt and registry snapshot without downloading or +reactivating the package. + +Publish metadata below `/metadata/` and payloads below +`/targets/`. An extension target uses this canonical path: + +```text +extensions////// +``` + +Its TUF target `custom.a3s` object must contain `schemaVersion`, `packageId`, +`version`, `channel` (`stable`, `beta`, or `nightly`), and `target` (an A3S host +target or `any`). Duplicate identities, mismatched paths, unsupported archives, +and oversized targets are rejected before payload download. Built-in and management routes are reserved. Extensions cannot shadow `browser`, `office`, `ocr`, `box`, `component`, `capability`, or other host diff --git a/crates/extension/Cargo.toml b/crates/extension/Cargo.toml index 93f63891..d1ea839f 100644 --- a/crates/extension/Cargo.toml +++ b/crates/extension/Cargo.toml @@ -12,9 +12,19 @@ description = "ACL manifest and native surface contracts for A3S Use extensions" a3s-acl = { git = "https://github.com/A3S-Lab/ACL", rev = "6e2a6469edc0f4c61b1e588d0ace873aaf15ce22" } a3s-use-core = { version = "0.1.1", path = "../core" } fs2.workspace = true +flate2.workspace = true +reqwest.workspace = true serde.workspace = true serde_json.workspace = true semver = "1" sha2.workspace = true +tar.workspace = true tempfile.workspace = true tokio.workspace = true +tough.workspace = true +url.workspace = true +zip.workspace = true + +[dev-dependencies] +olpc-cjson = "0.1" +ring = "0.17" diff --git a/crates/extension/src/digest.rs b/crates/extension/src/digest.rs new file mode 100644 index 00000000..7f8fb87c --- /dev/null +++ b/crates/extension/src/digest.rs @@ -0,0 +1,194 @@ +use std::fs::File; +use std::io::{BufReader, Read}; +use std::path::{Path, PathBuf}; + +use a3s_use_core::{UseError, UseResult}; +use sha2::{Digest, Sha256}; + +use super::package::{io_error, MAX_PACKAGE_BYTES, MAX_PACKAGE_FILES}; +use super::source::sanitized_relative_path; + +struct PackageFile { + normalized: String, + path: PathBuf, + size: u64, +} + +pub(crate) async fn package_sha256(root: &Path) -> UseResult { + let root = root.to_path_buf(); + tokio::task::spawn_blocking(move || hash_package(&root)) + .await + .map_err(|error| { + UseError::new( + "use.extension.io", + format!("Failed to hash extension package: blocking task failed: {error}"), + ) + })? +} + +fn hash_package(root: &Path) -> UseResult { + let mut files = Vec::new(); + let mut entries = 0_usize; + let mut bytes = 0_u64; + collect_files(root, root, &mut files, &mut entries, &mut bytes)?; + files.sort_by(|left, right| left.normalized.cmp(&right.normalized)); + + let mut digest = Sha256::new(); + digest.update(b"a3s-use-expanded-package-v1\0"); + for package_file in files { + let path_bytes = package_file.normalized.as_bytes(); + digest.update((path_bytes.len() as u64).to_be_bytes()); + digest.update(path_bytes); + digest.update(package_file.size.to_be_bytes()); + + let file = File::open(&package_file.path) + .map_err(|error| io_error("open extension package file", &package_file.path, error))?; + let mut reader = BufReader::new(file); + let mut buffer = [0_u8; 64 * 1024]; + let mut read_bytes = 0_u64; + loop { + let count = reader.read(&mut buffer).map_err(|error| { + io_error("hash extension package file", &package_file.path, error) + })?; + if count == 0 { + break; + } + read_bytes = read_bytes.saturating_add(count as u64); + if read_bytes > package_file.size { + return Err(package_changed(&package_file.path)); + } + digest.update(&buffer[..count]); + } + if read_bytes != package_file.size { + return Err(package_changed(&package_file.path)); + } + } + Ok(format!("{:x}", digest.finalize())) +} + +fn collect_files( + root: &Path, + directory: &Path, + files: &mut Vec, + entries: &mut usize, + bytes: &mut u64, +) -> UseResult<()> { + let children = std::fs::read_dir(directory) + .map_err(|error| io_error("read extension package directory", directory, error))?; + for child in children { + let child = + child.map_err(|error| io_error("read extension package entry", directory, error))?; + *entries = entries.saturating_add(1); + if *entries > MAX_PACKAGE_FILES { + return Err(package_limit_error()); + } + let path = child.path(); + let metadata = std::fs::symlink_metadata(&path) + .map_err(|error| io_error("inspect extension package entry", &path, error))?; + if metadata.file_type().is_symlink() { + return Err(UseError::new( + "use.extension.package_symlink", + format!( + "Extension package entry '{}' is a symbolic link.", + path.display() + ), + )); + } + if metadata.is_dir() { + collect_files(root, &path, files, entries, bytes)?; + continue; + } + if !metadata.is_file() { + return Err(UseError::new( + "use.extension.package_entry_invalid", + format!( + "Extension package entry '{}' is not a regular file or directory.", + path.display() + ), + )); + } + *bytes = bytes.saturating_add(metadata.len()); + if *bytes > MAX_PACKAGE_BYTES { + return Err(package_limit_error()); + } + let relative = path.strip_prefix(root).map_err(|_| { + UseError::new( + "use.extension.path_escape", + format!( + "Extension package entry '{}' escapes its root.", + path.display() + ), + ) + })?; + let relative = sanitized_relative_path(relative)?.ok_or_else(|| { + UseError::new( + "use.extension.package_entry_invalid", + "Extension package contains an empty file path.", + ) + })?; + let normalized = relative + .iter() + .map(|segment| { + segment.to_str().ok_or_else(|| { + UseError::new( + "use.extension.package_entry_invalid", + format!( + "Extension package path '{}' is not valid UTF-8.", + relative.display() + ), + ) + }) + }) + .collect::>>()? + .join("/"); + files.push(PackageFile { + normalized, + path, + size: metadata.len(), + }); + } + Ok(()) +} + +fn package_changed(path: &Path) -> UseError { + UseError::new( + "use.extension.package_changed", + format!( + "Extension package file '{}' changed while it was hashed.", + path.display() + ), + ) +} + +fn package_limit_error() -> UseError { + UseError::new( + "use.extension.package_too_large", + "The extension package exceeds the local installation limits.", + ) +} + +#[cfg(test)] +mod tests { + use super::*; + + #[tokio::test] + async fn package_digest_is_order_independent_and_content_sensitive() { + let temp = tempfile::tempdir().unwrap(); + let first = temp.path().join("first"); + let second = temp.path().join("second"); + std::fs::create_dir_all(first.join("bin")).unwrap(); + std::fs::create_dir_all(second.join("bin")).unwrap(); + std::fs::write(first.join("z.txt"), b"z").unwrap(); + std::fs::write(first.join("bin/tool"), b"tool").unwrap(); + std::fs::write(second.join("bin/tool"), b"tool").unwrap(); + std::fs::write(second.join("z.txt"), b"z").unwrap(); + + let first_digest = package_sha256(&first).await.unwrap(); + let second_digest = package_sha256(&second).await.unwrap(); + assert_eq!(first_digest, second_digest); + assert_eq!(first_digest.len(), 64); + + std::fs::write(second.join("bin/tool"), b"changed").unwrap(); + assert_ne!(first_digest, package_sha256(&second).await.unwrap()); + } +} diff --git a/crates/extension/src/lib.rs b/crates/extension/src/lib.rs index cc76e2dc..9b692c3c 100644 --- a/crates/extension/src/lib.rs +++ b/crates/extension/src/lib.rs @@ -5,11 +5,14 @@ use a3s_acl::{Block, Value}; use a3s_use_core::{RiskClass, UseError, UseResult}; use serde::{Deserialize, Serialize}; +mod digest; mod package; mod paths; mod registry; mod registry_io; +mod remote; mod route_lock; +mod source; pub use paths::ExtensionPaths; pub use registry::{ @@ -17,6 +20,10 @@ pub use registry::{ ExtensionRouteBinding, ExtensionRouteLease, ExtensionTrust, InstallOptions, InstallResult, InstalledExtension, UninstallResult, }; +pub use remote::{ + prepare_remote_package, refresh_remote_registry, DownloadedRemotePackage, + PreparedRemotePackage, ResolvedRemotePackage, TrustedRegistry, VerifiedRegistryMetadata, +}; const RESERVED_ROUTES: &[&str] = &[ "browser", diff --git a/crates/extension/src/package.rs b/crates/extension/src/package.rs index a939a430..9778ace7 100644 --- a/crates/extension/src/package.rs +++ b/crates/extension/src/package.rs @@ -12,9 +12,9 @@ use tokio::io::AsyncWriteExt; use super::registry::ExtensionReceipt; use super::{ExtensionManifest, ExtensionPaths}; -const MANIFEST_NAME: &str = "a3s-use-extension.acl"; -const MAX_PACKAGE_FILES: usize = 10_000; -const MAX_PACKAGE_BYTES: u64 = 1_073_741_824; +pub(crate) const MANIFEST_NAME: &str = "a3s-use-extension.acl"; +pub(crate) const MAX_PACKAGE_FILES: usize = 10_000; +pub(crate) const MAX_PACKAGE_BYTES: u64 = 1_073_741_824; pub(crate) async fn read_manifest(package_root: &Path) -> UseResult<(ExtensionManifest, Vec)> { let path = package_root.join(MANIFEST_NAME); diff --git a/crates/extension/src/paths.rs b/crates/extension/src/paths.rs index 62adef50..3232a610 100644 --- a/crates/extension/src/paths.rs +++ b/crates/extension/src/paths.rs @@ -93,6 +93,12 @@ impl ExtensionPaths { path.set_extension("lock"); path } + + pub fn tuf_datastore(&self, registry_name: &str) -> PathBuf { + self.state_root + .join("remote-registries") + .join(registry_name) + } } fn configured_root( @@ -163,5 +169,9 @@ mod tests { paths.registry_snapshot_path(), PathBuf::from("/state/use/registry.json") ); + assert_eq!( + paths.tuf_datastore("a3s"), + PathBuf::from("/state/use/remote-registries/a3s") + ); } } diff --git a/crates/extension/src/registry.rs b/crates/extension/src/registry.rs index 5f3c0135..05e3dc97 100644 --- a/crates/extension/src/registry.rs +++ b/crates/extension/src/registry.rs @@ -7,12 +7,15 @@ use fs2::FileExt; use serde::{Deserialize, Serialize}; use tokio::fs; +use super::digest::package_sha256; use super::package::{ copy_package, io_error, owned_package_path, read_manifest, sha256, unique_suffix, unix_timestamp, validate_surface_files, write_receipt, RegistryLock, }; use super::registry_io::{read_registry_snapshot, write_registry_snapshot}; +use super::remote::{prepare_remote_package, ResolvedRemotePackage, TrustedRegistry}; use super::route_lock::{acquire_drain_lock, deadline_after, open_route_lock}; +use super::source::prepare_package_source; use super::{ExtensionManifest, ExtensionPaths, McpTransport}; const RECEIPT_SCHEMA_VERSION: u32 = 1; @@ -24,6 +27,7 @@ const WATCH_INTERVAL: Duration = Duration::from_millis(50); #[serde(rename_all = "kebab-case")] pub enum ExtensionTrust { LocalExplicit, + RegistryTuf, } #[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] @@ -36,7 +40,11 @@ pub struct ExtensionReceipt { pub version: String, pub package_root: PathBuf, pub manifest_sha256: String, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub package_sha256: Option, pub trust: ExtensionTrust, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub registry: Option, pub installed_at_unix: u64, #[serde(default = "enabled_by_default")] pub enabled: bool, @@ -115,6 +123,8 @@ pub struct ExtensionRouteBinding { #[serde(default)] pub package_root: PathBuf, pub manifest_sha256: String, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub package_sha256: Option, pub enabled: bool, pub surfaces: Vec, } @@ -368,21 +378,113 @@ impl ExtensionRegistry { .with_suggestion("Rerun the explicit install with --allow-unsigned.")); } - let source = fs::canonicalize(source) - .await - .map_err(|error| io_error("resolve extension package", source, error))?; - let source_metadata = fs::metadata(&source) - .await - .map_err(|error| io_error("inspect extension package", &source, error))?; - if !source_metadata.is_dir() { - return Err(UseError::new( - "use.extension.package_unsupported", - "The initial local installer accepts a package directory.", + let source = prepare_package_source(source).await?; + self.install_prepared( + &expected_package_id, + source.root(), + options.force, + ExtensionTrust::LocalExplicit, + None, + ) + .await + } + + /// Install an extension selected through a fully verified TUF repository. + /// + /// Metadata is resolved and the optional reviewed plan is checked before + /// the target payload is downloaded. The package manifest must repeat the + /// exact ID and version carried by the signed target metadata. + pub async fn install_remote( + &self, + expected_package_id: &str, + registry: &TrustedRegistry, + requested_version: Option<&str>, + channel: &str, + expected_plan_digest: Option<&str>, + force: bool, + ) -> UseResult { + let expected_package_id = normalize_package_id(expected_package_id)?; + let prepared = prepare_remote_package( + registry, + &expected_package_id, + requested_version, + channel, + expected_plan_digest, + ) + .await?; + if !force { + if let Some(result) = self + .converged_remote_install(&expected_package_id, prepared.resolved()) + .await? + { + return Ok(result); + } + } + let downloaded = prepared.download().await?; + let provenance = downloaded.resolved().clone(); + let source = prepare_package_source(downloaded.path()).await?; + self.install_prepared( + &expected_package_id, + source.root(), + force, + ExtensionTrust::RegistryTuf, + Some(provenance), + ) + .await + } + + async fn converged_remote_install( + &self, + expected_package_id: &str, + resolved: &ResolvedRemotePackage, + ) -> UseResult> { + let _lock = RegistryLock::acquire(&self.paths.registry_lock_path())?; + let Some(mut current) = self.get(expected_package_id).await? else { + return Ok(None); + }; + let same_target = current.receipt.trust == ExtensionTrust::RegistryTuf + && current.receipt.version == resolved.version + && registry_identity(current.receipt.registry.as_ref()) + == registry_identity(Some(resolved)); + if !same_target { + return Ok(None); + } + verify_package_integrity(¤t).await?; + if current.receipt.registry.as_ref() != Some(resolved) { + current.receipt.registry = Some(resolved.clone()); + write_receipt( + &self.paths.receipt_path(expected_package_id), + ¤t.receipt, ) - .with_suggestion("Extract the package archive and pass its directory with --from.")); + .await?; + } + let installed = self.list().await?; + self.publish_snapshot_locked(&installed).await?; + Ok(Some(InstallResult { + changed: false, + extension: current, + })) + } + + async fn install_prepared( + &self, + expected_package_id: &str, + source: &Path, + force: bool, + trust: ExtensionTrust, + registry: Option, + ) -> UseResult { + match (trust, registry.as_ref()) { + (ExtensionTrust::LocalExplicit, None) | (ExtensionTrust::RegistryTuf, Some(_)) => {} + _ => { + return Err(UseError::new( + "use.extension.trust_invalid", + "Extension installation provenance is internally inconsistent.", + )) + } } - let (manifest, manifest_bytes) = read_manifest(&source).await?; + let (manifest, manifest_bytes) = read_manifest(source).await?; if manifest.package_id != expected_package_id { return Err(UseError::new( "use.extension.identity_mismatch", @@ -392,7 +494,22 @@ impl ExtensionRegistry { ), )); } - validate_surface_files(&manifest, &source).await?; + if let Some(registry) = ®istry { + if registry.package_id != manifest.package_id || registry.version != manifest.version { + return Err(UseError::new( + "use.extension.registry_identity_mismatch", + format!( + "Signed target '{}@{}' does not match package manifest '{}@{}'.", + registry.package_id, + registry.version, + manifest.package_id, + manifest.version + ), + )); + } + } + validate_surface_files(&manifest, source).await?; + let package_digest = package_sha256(source).await?; let _lock = RegistryLock::acquire(&self.paths.registry_lock_path())?; let installed = self.list().await?; @@ -414,9 +531,17 @@ impl ExtensionRegistry { .iter() .find(|extension| extension.receipt.package_id == expected_package_id) { - if !options.force + let current_package_digest = match ¤t.receipt.package_sha256 { + Some(digest) => digest.clone(), + None => package_sha256(¤t.receipt.package_root).await?, + }; + let same_provenance = current.receipt.trust == trust + && registry_identity(current.receipt.registry.as_ref()) + == registry_identity(registry.as_ref()); + if !force && current.receipt.version == manifest.version - && current.receipt.manifest_sha256 == digest + && current_package_digest == package_digest + && same_provenance { self.publish_snapshot_locked(&installed).await?; return Ok(InstallResult { @@ -424,7 +549,10 @@ impl ExtensionRegistry { extension: current.clone(), }); } - if !options.force && current.receipt.version == manifest.version { + if !force + && current.receipt.version == manifest.version + && current_package_digest != package_digest + { return Err(UseError::new( "use.extension.version_conflict", format!( @@ -436,7 +564,7 @@ impl ExtensionRegistry { } } - let package_parent = self.paths.package_parent(&expected_package_id); + let package_parent = self.paths.package_parent(expected_package_id); fs::create_dir_all(&package_parent).await.map_err(|error| { io_error("create extension package directory", &package_parent, error) })?; @@ -446,7 +574,7 @@ impl ExtensionRegistry { .map_err(|error| { io_error("create extension staging directory", &package_parent, error) })?; - copy_package(&source, staging.path()).await?; + copy_package(source, staging.path()).await?; let (staged_manifest, staged_bytes) = read_manifest(staging.path()).await?; if staged_manifest != manifest || sha256(&staged_bytes) != digest { return Err(UseError::new( @@ -455,11 +583,17 @@ impl ExtensionRegistry { )); } validate_surface_files(&staged_manifest, staging.path()).await?; + if package_sha256(staging.path()).await? != package_digest { + return Err(UseError::new( + "use.extension.package_changed", + "The extension package changed while it was staged.", + )); + } let activation = unique_suffix(); let target = self .paths - .package_root(&expected_package_id, &manifest.version, &activation); + .package_root(expected_package_id, &manifest.version, &activation); let staging = staging.keep(); if let Err(error) = fs::rename(&staging, &target).await { let _ = fs::remove_dir_all(&staging).await; @@ -474,17 +608,19 @@ impl ExtensionRegistry { let receipt = ExtensionReceipt { schema_version: RECEIPT_SCHEMA_VERSION, - package_id: expected_package_id.clone(), + package_id: expected_package_id.to_string(), component_id: format!("use/{expected_package_id}"), route: manifest.route.clone(), version: manifest.version.clone(), package_root: target.clone(), manifest_sha256: digest, - trust: ExtensionTrust::LocalExplicit, + package_sha256: Some(package_digest), + trust, + registry, installed_at_unix: unix_timestamp(), enabled, }; - let receipt_path = self.paths.receipt_path(&expected_package_id); + let receipt_path = self.paths.receipt_path(expected_package_id); if let Err(error) = write_receipt(&receipt_path, &receipt).await { let _ = fs::remove_dir_all(&target).await; return Err(error); @@ -650,6 +786,7 @@ impl ExtensionRegistry { let _ = FileExt::unlock(&file); return Ok(None); } + verify_package_integrity(&extension).await?; Ok(Some(ExtensionRouteLease { extension, file })) } @@ -699,6 +836,46 @@ impl ExtensionRegistry { ), )); } + if receipt.package_sha256.as_deref().is_some_and(|digest| { + digest.len() != 64 || !digest.bytes().all(|byte| byte.is_ascii_hexdigit()) + }) { + return Err(UseError::new( + "use.extension.receipt_invalid", + format!( + "Extension receipt for '{}' has an invalid package digest.", + receipt.package_id + ), + )); + } + match ( + receipt.trust, + receipt.registry.as_ref(), + receipt.package_sha256.as_ref(), + ) { + (ExtensionTrust::LocalExplicit, None, _) => {} + (ExtensionTrust::RegistryTuf, Some(registry), Some(_)) => { + registry.validate_provenance()?; + if registry.package_id != receipt.package_id || registry.version != receipt.version + { + return Err(UseError::new( + "use.extension.receipt_invalid", + format!( + "Registry provenance for '{}' does not match its receipt.", + receipt.package_id + ), + )); + } + } + _ => { + return Err(UseError::new( + "use.extension.receipt_invalid", + format!( + "Extension receipt for '{}' has inconsistent trust provenance.", + receipt.package_id + ), + )) + } + } let package_id = normalize_package_id(&receipt.package_id)?; if receipt.component_id != format!("use/{package_id}") || !owned_package_path(&self.paths, &package_id, &receipt.package_root) @@ -730,6 +907,24 @@ impl ExtensionRegistry { } } +async fn verify_package_integrity(extension: &InstalledExtension) -> UseResult<()> { + let Some(expected) = extension.receipt.package_sha256.as_deref() else { + return Ok(()); + }; + let actual = package_sha256(&extension.receipt.package_root).await?; + if actual != expected { + return Err(UseError::new( + "use.extension.package_digest_mismatch", + format!( + "Installed package '{}' no longer matches its recorded digest.", + extension.receipt.package_id + ), + ) + .with_suggestion("Reinstall the extension from its trusted source.")); + } + Ok(()) +} + fn route_bindings(installed: &[InstalledExtension]) -> Vec { installed .iter() @@ -740,6 +935,7 @@ fn route_bindings(installed: &[InstalledExtension]) -> Vec UseResult { Ok(value.to_string()) } +fn registry_identity(registry: Option<&ResolvedRemotePackage>) -> Option<(&str, &str, &str, &str)> { + registry.map(|registry| { + ( + registry.registry_name.as_str(), + registry.registry_url.as_str(), + registry.root_sha256.as_str(), + registry.sha256.as_str(), + ) + }) +} + fn ensure_unique_routes(installed: &[InstalledExtension]) -> UseResult<()> { for (index, extension) in installed.iter().enumerate() { if let Some(conflict) = installed[index + 1..] diff --git a/crates/extension/src/registry_tests.rs b/crates/extension/src/registry_tests.rs index 59664552..875fd2bc 100644 --- a/crates/extension/src/registry_tests.rs +++ b/crates/extension/src/registry_tests.rs @@ -1,3 +1,5 @@ +use std::fs::File; +use std::io::Write; use std::time::Duration; #[cfg(unix)] @@ -48,6 +50,43 @@ fn registry(root: &Path) -> ExtensionRegistry { ExtensionRegistry::new(ExtensionPaths::new(root.join("data"), root.join("state"))) } +fn tar_package(source: &Path, archive: &Path) { + let file = File::create(archive).unwrap(); + let encoder = flate2::write::GzEncoder::new(file, flate2::Compression::default()); + let mut builder = tar::Builder::new(encoder); + builder.append_dir_all("package", source).unwrap(); + builder.finish().unwrap(); +} + +fn zip_package(source: &Path, archive: &Path) { + let file = File::create(archive).unwrap(); + let mut writer = zip::ZipWriter::new(file); + for relative in [ + "a3s-use-extension.acl", + "bin/extension", + "skills/demo/SKILL.md", + ] { + let source_file = source.join(relative); + let mut options = zip::write::SimpleFileOptions::default() + .compression_method(zip::CompressionMethod::Deflated); + #[cfg(unix)] + { + let mode = std::fs::metadata(&source_file) + .unwrap() + .permissions() + .mode(); + options = options.unix_permissions(mode); + } + writer + .start_file(format!("package/{relative}"), options) + .unwrap(); + writer + .write_all(&std::fs::read(source_file).unwrap()) + .unwrap(); + } + writer.finish().unwrap(); +} + #[tokio::test] async fn installs_lists_and_uninstalls_an_explicit_local_package() { let temp = tempfile::tempdir().unwrap(); @@ -89,6 +128,63 @@ async fn installs_lists_and_uninstalls_an_explicit_local_package() { assert!(registry.list().await.unwrap().is_empty()); } +#[tokio::test] +async fn installs_and_uninstalls_a_local_tar_package() { + let temp = tempfile::tempdir().unwrap(); + let source = temp.path().join("source"); + package(&source, "acme/slack", "slack", "1.2.0").await; + let archive = temp.path().join("acme-slack.tar.gz"); + tar_package(&source, &archive); + let registry = registry(temp.path()); + + let result = registry + .install_local( + "acme/slack", + &archive, + InstallOptions { + allow_unsigned: true, + force: false, + }, + ) + .await + .unwrap(); + assert!(result.changed); + assert_eq!(result.extension.receipt.package_id, "acme/slack"); + assert!(result.extension.cli_executable().unwrap().is_file()); + + let removed = registry.uninstall("acme/slack").await.unwrap(); + assert!(removed.changed); + assert!(registry.list().await.unwrap().is_empty()); +} + +#[tokio::test] +async fn installs_and_uninstalls_a_local_zip_package() { + let temp = tempfile::tempdir().unwrap(); + let source = temp.path().join("source"); + package(&source, "acme/slack", "slack", "1.2.0").await; + let archive = temp.path().join("acme-slack.zip"); + zip_package(&source, &archive); + let registry = registry(temp.path()); + + let result = registry + .install_local( + "acme/slack", + &archive, + InstallOptions { + allow_unsigned: true, + force: false, + }, + ) + .await + .unwrap(); + assert!(result.changed); + assert_eq!(result.extension.receipt.package_id, "acme/slack"); + assert!(result.extension.cli_executable().unwrap().is_file()); + + assert!(registry.uninstall("acme/slack").await.unwrap().changed); + assert!(registry.list().await.unwrap().is_empty()); +} + #[tokio::test] async fn rejects_route_conflicts_and_untrusted_installs() { let temp = tempfile::tempdir().unwrap(); @@ -198,6 +294,8 @@ async fn hot_upgrade_keeps_the_previous_package_until_inflight_routes_drain() { let second = temp.path().join("second"); package(&first, "acme/slack", "slack", "1.0.0").await; package(&second, "acme/slack", "slack", "2.0.0").await; + let second_archive = temp.path().join("second.tar.gz"); + tar_package(&second, &second_archive); let registry = registry(temp.path()); let first_install = registry @@ -217,7 +315,7 @@ async fn hot_upgrade_keeps_the_previous_package_until_inflight_routes_drain() { let second_install = registry .install_local( "acme/slack", - &second, + &second_archive, InstallOptions { allow_unsigned: true, force: false, @@ -272,6 +370,16 @@ async fn forced_reactivation_of_identical_metadata_publishes_a_new_generation() second.extension.receipt.package_root, first.extension.receipt.package_root ); + assert_eq!( + second.extension.receipt.package_sha256, + first.extension.receipt.package_sha256 + ); + assert!(second + .extension + .receipt + .package_sha256 + .as_deref() + .is_some_and(|digest| digest.len() == 64)); let second_snapshot = registry.snapshot().await.unwrap(); assert_eq!(second_snapshot.generation, 2); assert_eq!( @@ -280,6 +388,148 @@ async fn forced_reactivation_of_identical_metadata_publishes_a_new_generation() ); } +#[tokio::test] +async fn same_version_changed_executable_requires_force_and_changes_package_digest() { + let temp = tempfile::tempdir().unwrap(); + let source = temp.path().join("source"); + package(&source, "acme/slack", "slack", "1.0.0").await; + let registry = registry(temp.path()); + + let first = registry + .install_local( + "acme/slack", + &source, + InstallOptions { + allow_unsigned: true, + force: false, + }, + ) + .await + .unwrap(); + fs::write( + source.join("bin/extension"), + "#!/bin/sh\nprintf 'changed\\n'\n", + ) + .await + .unwrap(); + + let error = registry + .install_local( + "acme/slack", + &source, + InstallOptions { + allow_unsigned: true, + force: false, + }, + ) + .await + .unwrap_err(); + assert_eq!(error.code, "use.extension.version_conflict"); + + let second = registry + .install_local( + "acme/slack", + &source, + InstallOptions { + allow_unsigned: true, + force: true, + }, + ) + .await + .unwrap(); + assert_ne!( + second.extension.receipt.package_root, + first.extension.receipt.package_root + ); + assert_ne!( + second.extension.receipt.package_sha256, + first.extension.receipt.package_sha256 + ); + assert!(second.extension.receipt.package_sha256.is_some()); + assert_eq!( + fs::read_to_string(second.extension.cli_executable().unwrap()) + .await + .unwrap(), + "#!/bin/sh\nprintf 'changed\\n'\n" + ); +} + +#[tokio::test] +async fn legacy_receipt_without_package_digest_remains_readable_and_idempotent() { + let temp = tempfile::tempdir().unwrap(); + let source = temp.path().join("source"); + package(&source, "acme/slack", "slack", "1.0.0").await; + let registry = registry(temp.path()); + + registry + .install_local( + "acme/slack", + &source, + InstallOptions { + allow_unsigned: true, + force: false, + }, + ) + .await + .unwrap(); + + let receipt_path = registry.paths().receipt_path("acme/slack"); + let mut legacy: serde_json::Value = + serde_json::from_slice(&fs::read(&receipt_path).await.unwrap()).unwrap(); + legacy.as_object_mut().unwrap().remove("packageSha256"); + fs::write(&receipt_path, serde_json::to_vec_pretty(&legacy).unwrap()) + .await + .unwrap(); + + let installed = registry.get("acme/slack").await.unwrap().unwrap(); + assert_eq!(installed.receipt.package_sha256, None); + + let unchanged = registry + .install_local( + "acme/slack", + &source, + InstallOptions { + allow_unsigned: true, + force: false, + }, + ) + .await + .unwrap(); + assert!(!unchanged.changed); + assert_eq!(unchanged.extension.receipt.package_sha256, None); +} + +#[tokio::test] +async fn receipt_rejects_an_invalid_optional_package_digest() { + let temp = tempfile::tempdir().unwrap(); + let source = temp.path().join("source"); + package(&source, "acme/slack", "slack", "1.0.0").await; + let registry = registry(temp.path()); + + registry + .install_local( + "acme/slack", + &source, + InstallOptions { + allow_unsigned: true, + force: false, + }, + ) + .await + .unwrap(); + + let receipt_path = registry.paths().receipt_path("acme/slack"); + let mut invalid: serde_json::Value = + serde_json::from_slice(&fs::read(&receipt_path).await.unwrap()).unwrap(); + invalid["packageSha256"] = serde_json::json!("not-a-sha256"); + fs::write(&receipt_path, serde_json::to_vec_pretty(&invalid).unwrap()) + .await + .unwrap(); + + let error = registry.get("acme/slack").await.unwrap_err(); + assert_eq!(error.code, "use.extension.receipt_invalid"); +} + #[tokio::test] async fn snapshot_reconciles_a_pre_activation_identity_binding() { let temp = tempfile::tempdir().unwrap(); diff --git a/crates/extension/src/remote.rs b/crates/extension/src/remote.rs new file mode 100644 index 00000000..8bd205b1 --- /dev/null +++ b/crates/extension/src/remote.rs @@ -0,0 +1,970 @@ +//! TUF-backed remote extension registry resolution. +//! +//! The trusted root is pinned out of band by SHA-256. Tough then verifies the +//! complete root/timestamp/snapshot/targets chain, enforces expiration, and +//! persists metadata versions in its datastore to reject rollback attacks. + +use std::collections::BTreeSet; +use std::fs::{File, OpenOptions}; +use std::path::{Path, PathBuf}; +use std::time::Duration; + +use a3s_use_core::{UseError, UseResult}; +use fs2::FileExt; +use semver::Version; +use serde::{Deserialize, Serialize}; +use sha2::{Digest, Sha256}; +use tempfile::TempDir; +use tokio::fs; +use tokio::io::AsyncWriteExt; +use tough::{ExpirationEnforcement, HttpTransportBuilder, Limits, Prefix, Repository}; +use tough::{RepositoryLoader, TargetName}; +use url::Url; + +use super::package::{activate_temporary_file, io_error, sync_parent_directory, unique_suffix}; + +const ROOT_NAME: &str = "root.json"; +const ROOT_CACHE_NAME: &str = "bootstrap-root.json"; +const REGISTRY_METADATA_KEY: &str = "a3s"; +const REGISTRY_TARGET_SCHEMA_VERSION: u32 = 1; +const MAX_BOOTSTRAP_ROOT_BYTES: u64 = 1024 * 1024; +const MAX_REMOTE_ARCHIVE_BYTES: u64 = 512 * 1024 * 1024; +const MAX_ROOT_UPDATES: u64 = 64; + +/// One configured registry whose TUF root is pinned out of band. +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct TrustedRegistry { + name: String, + base_url: Url, + root_sha256: String, + trusted_root_path: Option, + datastore: PathBuf, +} + +impl TrustedRegistry { + pub fn new( + name: impl Into, + base_url: impl AsRef, + root_sha256: impl AsRef, + trusted_root_path: Option, + datastore: PathBuf, + ) -> UseResult { + let name = name.into(); + validate_registry_name(&name)?; + let base_url = normalize_registry_url(base_url.as_ref())?; + let root_sha256 = normalize_sha256(root_sha256.as_ref(), "registry trust root")?; + if !datastore.is_absolute() { + return Err(UseError::new( + "use.extension.registry_path_invalid", + "The TUF metadata datastore must be an absolute path.", + )); + } + if trusted_root_path + .as_ref() + .is_some_and(|path| !path.is_absolute()) + { + return Err(UseError::new( + "use.extension.registry_path_invalid", + "The trusted TUF root path must be absolute.", + )); + } + Ok(Self { + name, + base_url, + root_sha256, + trusted_root_path, + datastore, + }) + } + + pub fn name(&self) -> &str { + &self.name + } + + pub fn base_url(&self) -> &Url { + &self.base_url + } + + pub fn root_sha256(&self) -> &str { + &self.root_sha256 + } + + pub fn datastore(&self) -> &Path { + &self.datastore + } + + fn metadata_url(&self) -> UseResult { + self.base_url.join("metadata/").map_err(|error| { + UseError::new( + "use.extension.registry_url_invalid", + format!("Failed to resolve the registry metadata URL: {error}"), + ) + }) + } + + fn targets_url(&self) -> UseResult { + self.base_url.join("targets/").map_err(|error| { + UseError::new( + "use.extension.registry_url_invalid", + format!("Failed to resolve the registry targets URL: {error}"), + ) + }) + } +} + +/// Exact signed target selected from a verified TUF repository. +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +#[serde(rename_all = "camelCase")] +pub struct ResolvedRemotePackage { + pub registry_name: String, + pub registry_url: String, + pub root_sha256: String, + pub root_version: u64, + pub timestamp_version: u64, + pub snapshot_version: u64, + pub targets_version: u64, + pub package_id: String, + pub version: String, + pub channel: String, + pub target: String, + pub target_name: String, + pub archive_name: String, + pub length: u64, + pub sha256: String, +} + +/// Signed metadata versions observed after a complete TUF refresh. +#[derive(Debug, Clone, PartialEq, Eq, Serialize)] +#[serde(rename_all = "camelCase")] +pub struct VerifiedRegistryMetadata { + pub registry_name: String, + pub registry_url: String, + pub root_sha256: String, + pub root_version: u64, + pub timestamp_version: u64, + pub snapshot_version: u64, + pub targets_version: u64, + pub package_targets: u64, +} + +impl ResolvedRemotePackage { + pub fn plan_digest(&self) -> UseResult { + let bytes = serde_json::to_vec(self).map_err(|error| { + UseError::new( + "use.extension.registry_plan_invalid", + format!("Failed to encode the resolved registry plan: {error}"), + ) + })?; + Ok(format!("{:x}", Sha256::digest(bytes))) + } + + pub fn verify_expected_plan(&self, expected: Option<&str>) -> UseResult<()> { + let Some(expected) = expected else { + return Ok(()); + }; + let expected = normalize_sha256(expected, "expected registry plan")?; + let actual = self.plan_digest()?; + if expected == actual { + return Ok(()); + } + Err(UseError::new( + "use.extension.registry_plan_mismatch", + "The signed registry target changed after review.", + ) + .with_detail("expected", expected) + .with_detail("actual", actual)) + } + + pub(crate) fn validate_provenance(&self) -> UseResult<()> { + validate_registry_name(&self.registry_name)?; + let normalized_url = normalize_registry_url(&self.registry_url)?; + if normalized_url.as_str() != self.registry_url { + return Err(UseError::new( + "use.extension.receipt_invalid", + "The registry URL in the extension receipt is not canonical.", + )); + } + normalize_sha256(&self.root_sha256, "registry trust root")?; + normalize_sha256(&self.sha256, "registry target")?; + if self.root_version == 0 + || self.timestamp_version == 0 + || self.snapshot_version == 0 + || self.targets_version == 0 + || self.length == 0 + || self.length > MAX_REMOTE_ARCHIVE_BYTES + || !super::valid_package_id(&self.package_id) + || Version::parse(&self.version).is_err() + { + return Err(UseError::new( + "use.extension.receipt_invalid", + "The registry provenance in the extension receipt is invalid.", + )); + } + validate_channel(&self.channel)?; + let host = host_target()?; + if self.target != host && self.target != "any" { + return Err(UseError::new( + "use.extension.receipt_invalid", + "The installed registry target does not match this platform.", + )); + } + let target_name = TargetName::new(self.target_name.clone()).map_err(|error| { + UseError::new( + "use.extension.receipt_invalid", + format!("The registry target name in the receipt is invalid: {error}"), + ) + })?; + validate_target_name( + &target_name, + &RegistryTargetMetadata { + schema_version: REGISTRY_TARGET_SCHEMA_VERSION, + package_id: self.package_id.clone(), + version: self.version.clone(), + channel: self.channel.clone(), + target: self.target.clone(), + }, + )?; + if target_name.raw().rsplit('/').next() != Some(self.archive_name.as_str()) { + return Err(UseError::new( + "use.extension.receipt_invalid", + "The registry archive name does not match its signed target path.", + )); + } + Ok(()) + } +} + +/// Verified repository state retained until its exact target is downloaded. +pub struct PreparedRemotePackage { + repository: Repository, + target_name: TargetName, + resolved: ResolvedRemotePackage, +} + +impl std::fmt::Debug for PreparedRemotePackage { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("PreparedRemotePackage") + .field("resolved", &self.resolved) + .finish_non_exhaustive() + } +} + +impl PreparedRemotePackage { + pub fn resolved(&self) -> &ResolvedRemotePackage { + &self.resolved + } + + pub async fn download(self) -> UseResult { + let temporary = tokio::task::spawn_blocking(tempfile::tempdir) + .await + .map_err(|error| { + UseError::new( + "use.extension.registry_download_failed", + format!("Failed to create the remote package staging task: {error}"), + ) + })? + .map_err(|error| { + UseError::new( + "use.extension.registry_download_failed", + format!("Failed to create remote package staging: {error}"), + ) + })?; + self.repository + .save_target(&self.target_name, temporary.path(), Prefix::None) + .await + .map_err(|error| { + UseError::new( + "use.extension.registry_download_failed", + format!( + "Failed to download and verify TUF target '{}': {error}", + self.resolved.target_name + ), + ) + })?; + let path = temporary.path().join(self.target_name.resolved()); + let metadata = fs::metadata(&path) + .await + .map_err(|error| io_error("inspect downloaded TUF target", &path, error))?; + if !metadata.is_file() || metadata.len() != self.resolved.length { + return Err(UseError::new( + "use.extension.registry_target_invalid", + "The downloaded TUF target does not match its signed length.", + )); + } + Ok(DownloadedRemotePackage { + path, + resolved: self.resolved, + _temporary: temporary, + }) + } +} + +/// One downloaded archive kept alive through extension activation. +#[derive(Debug)] +pub struct DownloadedRemotePackage { + path: PathBuf, + resolved: ResolvedRemotePackage, + _temporary: TempDir, +} + +impl DownloadedRemotePackage { + pub fn path(&self) -> &Path { + &self.path + } + + pub fn resolved(&self) -> &ResolvedRemotePackage { + &self.resolved + } +} + +#[derive(Debug, Clone, Deserialize)] +#[serde(rename_all = "camelCase", deny_unknown_fields)] +struct RegistryTargetMetadata { + schema_version: u32, + package_id: String, + version: String, + channel: String, + target: String, +} + +struct MetadataLock(File); + +impl Drop for MetadataLock { + fn drop(&mut self) { + let _ = FileExt::unlock(&self.0); + } +} + +/// Load and verify a TUF repository, then select one exact extension target. +pub async fn prepare_remote_package( + registry: &TrustedRegistry, + package_id: &str, + requested_version: Option<&str>, + channel: &str, + expected_plan_digest: Option<&str>, +) -> UseResult { + if !super::valid_package_id(package_id) { + return Err(UseError::new( + "use.extension.id_invalid", + "Extension IDs must be '/' lowercase identifiers.", + )); + } + let requested_version = requested_version + .map(|version| { + Version::parse(version).map_err(|error| { + UseError::new( + "use.extension.version_invalid", + format!("Invalid requested extension version: {error}"), + ) + }) + }) + .transpose()?; + validate_channel(channel)?; + let repository = load_repository(registry).await?; + + let host_target = host_target()?; + let mut candidates = Vec::new(); + let mut identities = BTreeSet::new(); + for (target_name, target) in repository.all_targets() { + let Some(metadata) = target.custom.get(REGISTRY_METADATA_KEY) else { + continue; + }; + let metadata: RegistryTargetMetadata = + serde_json::from_value(metadata.clone()).map_err(|error| { + UseError::new( + "use.extension.registry_target_invalid", + format!( + "TUF target '{}' has invalid A3S metadata: {error}", + target_name.raw() + ), + ) + })?; + validate_target_metadata(target_name, target, &metadata)?; + let identity = ( + metadata.package_id.clone(), + metadata.version.clone(), + metadata.channel.clone(), + metadata.target.clone(), + ); + if !identities.insert(identity) { + return Err(UseError::new( + "use.extension.registry_target_invalid", + "The TUF repository contains duplicate A3S package targets.", + )); + } + if metadata.package_id != package_id + || metadata.channel != channel + || (metadata.target != host_target && metadata.target != "any") + { + continue; + } + let version = Version::parse(&metadata.version).map_err(|error| { + UseError::new( + "use.extension.registry_target_invalid", + format!( + "TUF target '{}' declares an invalid version: {error}", + target_name.raw() + ), + ) + })?; + if requested_version + .as_ref() + .is_some_and(|requested| requested != &version) + { + continue; + } + candidates.push((version, metadata, target_name.clone(), target.clone())); + } + candidates.sort_by(|left, right| { + left.0 + .cmp(&right.0) + .then_with(|| (left.1.target == host_target).cmp(&(right.1.target == host_target))) + .then_with(|| left.2.raw().cmp(right.2.raw())) + }); + let Some((version, metadata, target_name, target)) = candidates.pop() else { + return Err(UseError::new( + "use.extension.registry_package_missing", + format!( + "Registry '{}' has no '{}' package for channel '{}' and target '{}'.", + registry.name, package_id, channel, host_target + ), + )); + }; + if candidates.last().is_some_and(|candidate| { + candidate.0 == version + && (candidate.1.target == host_target) == (metadata.target == host_target) + }) { + return Err(UseError::new( + "use.extension.registry_target_invalid", + "The TUF repository resolves the same package version to multiple targets.", + )); + } + let archive_name = target_name + .raw() + .rsplit('/') + .next() + .unwrap_or_default() + .to_string(); + let resolved = ResolvedRemotePackage { + registry_name: registry.name.clone(), + registry_url: registry.base_url.to_string(), + root_sha256: registry.root_sha256.clone(), + root_version: repository.root().signed.version.get(), + timestamp_version: repository.timestamp().signed.version.get(), + snapshot_version: repository.snapshot().signed.version.get(), + targets_version: repository.targets().signed.version.get(), + package_id: package_id.to_string(), + version: version.to_string(), + channel: channel.to_string(), + target: metadata.target, + target_name: target_name.raw().to_string(), + archive_name, + length: target.length, + sha256: hex_lower(target.hashes.sha256.as_ref()), + }; + resolved.verify_expected_plan(expected_plan_digest)?; + Ok(PreparedRemotePackage { + repository, + target_name, + resolved, + }) +} + +/// Refresh and fully verify a registry without downloading any package target. +pub async fn refresh_remote_registry( + registry: &TrustedRegistry, +) -> UseResult { + let repository = load_repository(registry).await?; + let mut identities = BTreeSet::new(); + let mut package_targets = 0_u64; + for (target_name, target) in repository.all_targets() { + let Some(metadata) = target.custom.get(REGISTRY_METADATA_KEY) else { + continue; + }; + let metadata: RegistryTargetMetadata = + serde_json::from_value(metadata.clone()).map_err(|error| { + UseError::new( + "use.extension.registry_target_invalid", + format!( + "TUF target '{}' has invalid A3S metadata: {error}", + target_name.raw() + ), + ) + })?; + validate_target_metadata(target_name, target, &metadata)?; + let identity = ( + metadata.package_id, + metadata.version, + metadata.channel, + metadata.target, + ); + if !identities.insert(identity) { + return Err(UseError::new( + "use.extension.registry_target_invalid", + "The TUF repository contains duplicate A3S package targets.", + )); + } + package_targets = package_targets.checked_add(1).ok_or_else(|| { + UseError::new( + "use.extension.registry_target_invalid", + "The TUF repository contains too many package targets.", + ) + })?; + } + Ok(VerifiedRegistryMetadata { + registry_name: registry.name.clone(), + registry_url: registry.base_url.to_string(), + root_sha256: registry.root_sha256.clone(), + root_version: repository.root().signed.version.get(), + timestamp_version: repository.timestamp().signed.version.get(), + snapshot_version: repository.snapshot().signed.version.get(), + targets_version: repository.targets().signed.version.get(), + package_targets, + }) +} + +async fn load_repository(registry: &TrustedRegistry) -> UseResult { + ensure_metadata_directory(®istry.datastore).await?; + let lock = acquire_metadata_lock(®istry.datastore)?; + let root = load_trusted_root(registry).await?; + let metadata_url = registry.metadata_url()?; + let targets_url = registry.targets_url()?; + let transport = HttpTransportBuilder::new() + .timeout(Duration::from_secs(300)) + .connect_timeout(Duration::from_secs(15)) + .tries(3) + .build(); + let repository = RepositoryLoader::new(&root, metadata_url, targets_url) + .transport(transport) + .datastore(®istry.datastore) + .limits(Limits { + max_root_size: MAX_BOOTSTRAP_ROOT_BYTES, + max_targets_size: 10 * 1024 * 1024, + max_timestamp_size: 1024 * 1024, + max_snapshot_size: 1024 * 1024, + max_root_updates: MAX_ROOT_UPDATES, + }) + .expiration_enforcement(ExpirationEnforcement::Safe) + .load() + .await + .map_err(|error| { + UseError::new( + "use.extension.registry_untrusted", + format!( + "TUF verification failed for registry '{}': {error}", + registry.name + ), + ) + })?; + drop(lock); + Ok(repository) +} + +fn validate_target_metadata( + target_name: &TargetName, + target: &tough::schema::Target, + metadata: &RegistryTargetMetadata, +) -> UseResult<()> { + if metadata.schema_version != REGISTRY_TARGET_SCHEMA_VERSION { + return Err(UseError::new( + "use.extension.registry_target_invalid", + format!( + "TUF target '{}' uses unsupported A3S metadata schema {}.", + target_name.raw(), + metadata.schema_version + ), + )); + } + if !super::valid_package_id(&metadata.package_id) { + return Err(UseError::new( + "use.extension.registry_target_invalid", + format!( + "TUF target '{}' has an invalid package ID.", + target_name.raw() + ), + )); + } + Version::parse(&metadata.version).map_err(|error| { + UseError::new( + "use.extension.registry_target_invalid", + format!( + "TUF target '{}' has an invalid package version: {error}", + target_name.raw() + ), + ) + })?; + validate_channel(&metadata.channel)?; + validate_target_name(target_name, metadata)?; + if target.length == 0 || target.length > MAX_REMOTE_ARCHIVE_BYTES { + return Err(UseError::new( + "use.extension.registry_target_invalid", + format!( + "TUF target '{}' exceeds the supported package size.", + target_name.raw() + ), + )); + } + let digest = target.hashes.sha256.as_ref(); + if digest.len() != 32 { + return Err(UseError::new( + "use.extension.registry_target_invalid", + format!( + "TUF target '{}' does not have a valid SHA-256 digest.", + target_name.raw() + ), + )); + } + Ok(()) +} + +fn validate_target_name( + target_name: &TargetName, + metadata: &RegistryTargetMetadata, +) -> UseResult<()> { + let raw = target_name.raw(); + if raw != target_name.resolved() + || raw.starts_with('/') + || raw.contains('\\') + || raw.split('/').any(str::is_empty) + { + return Err(UseError::new( + "use.extension.registry_target_invalid", + format!("TUF target '{raw}' is not a portable package path."), + )); + } + let archive = raw.rsplit('/').next().unwrap_or_default(); + if !(archive.ends_with(".tar.gz") || archive.ends_with(".tgz") || archive.ends_with(".zip")) { + return Err(UseError::new( + "use.extension.registry_target_invalid", + format!("TUF target '{raw}' is not a supported package archive."), + )); + } + let expected_prefix = format!( + "extensions/{}/{}/{}/{}/", + metadata.package_id, metadata.version, metadata.channel, metadata.target + ); + if !raw.starts_with(&expected_prefix) { + return Err(UseError::new( + "use.extension.registry_target_invalid", + format!("TUF target '{raw}' must be published below '{expected_prefix}'."), + )); + } + Ok(()) +} + +fn validate_channel(channel: &str) -> UseResult<()> { + if matches!(channel, "stable" | "beta" | "nightly") { + Ok(()) + } else { + Err(UseError::new( + "use.extension.registry_channel_invalid", + format!("Unsupported extension release channel '{channel}'."), + )) + } +} + +fn host_target() -> UseResult { + match (std::env::consts::OS, std::env::consts::ARCH) { + ("macos", "aarch64") => Ok("darwin-arm64".to_string()), + ("macos", "x86_64") => Ok("darwin-x86_64".to_string()), + ("linux", "aarch64") => Ok("linux-arm64".to_string()), + ("linux", "x86_64") => Ok("linux-x86_64".to_string()), + ("windows", "x86_64") => Ok("windows-x86_64".to_string()), + (os, arch) => Err(UseError::new( + "use.extension.registry_target_unsupported", + format!("Remote extension packages are unavailable for {os}-{arch}."), + )), + } +} + +async fn ensure_metadata_directory(path: &Path) -> UseResult<()> { + fs::create_dir_all(path) + .await + .map_err(|error| io_error("create TUF metadata datastore", path, error))?; + let metadata = fs::symlink_metadata(path) + .await + .map_err(|error| io_error("inspect TUF metadata datastore", path, error))?; + if metadata.file_type().is_symlink() || !metadata.is_dir() { + return Err(UseError::new( + "use.extension.registry_path_invalid", + format!( + "The TUF metadata datastore '{}' must be a real directory.", + path.display() + ), + )); + } + #[cfg(unix)] + { + use std::os::unix::fs::PermissionsExt; + fs::set_permissions(path, std::fs::Permissions::from_mode(0o700)) + .await + .map_err(|error| io_error("secure TUF metadata datastore", path, error))?; + } + Ok(()) +} + +fn acquire_metadata_lock(datastore: &Path) -> UseResult { + let path = datastore.join(".metadata.lock"); + let file = OpenOptions::new() + .create(true) + .read(true) + .write(true) + .truncate(false) + .open(&path) + .map_err(|error| io_error("open TUF metadata lock", &path, error))?; + file.try_lock_exclusive().map_err(|error| { + UseError::new( + "use.extension.registry_busy", + format!( + "Another process is updating registry metadata '{}': {error}", + datastore.display() + ), + ) + })?; + Ok(MetadataLock(file)) +} + +async fn load_trusted_root(registry: &TrustedRegistry) -> UseResult> { + let explicit = registry.trusted_root_path.as_deref(); + let cache = registry.datastore.join(ROOT_CACHE_NAME); + let path = explicit.unwrap_or(&cache); + let bytes = match fs::read(path).await { + Ok(bytes) => bytes, + Err(error) if error.kind() == std::io::ErrorKind::NotFound && explicit.is_none() => { + let metadata_url = registry.metadata_url()?; + let root_url = metadata_url.join(ROOT_NAME).map_err(|error| { + UseError::new( + "use.extension.registry_url_invalid", + format!("Failed to resolve the bootstrap root URL: {error}"), + ) + })?; + let bytes = download_bootstrap_root(&root_url).await?; + verify_root_digest(registry, &bytes)?; + write_bootstrap_root(&cache, &bytes).await?; + bytes + } + Err(error) => return Err(io_error("read trusted TUF root", path, error)), + }; + if bytes.len() as u64 > MAX_BOOTSTRAP_ROOT_BYTES { + return Err(UseError::new( + "use.extension.registry_root_invalid", + "The trusted TUF root exceeds the one MiB limit.", + )); + } + verify_root_digest(registry, &bytes)?; + Ok(bytes) +} + +async fn download_bootstrap_root(url: &Url) -> UseResult> { + validate_download_url(url)?; + let client = reqwest::Client::builder() + .user_agent("a3s-use-extension/0.1") + .connect_timeout(Duration::from_secs(15)) + .timeout(Duration::from_secs(30)) + .redirect(reqwest::redirect::Policy::limited(5)) + .build() + .map_err(|error| { + UseError::new( + "use.extension.registry_download_failed", + format!("Failed to build the registry client: {error}"), + ) + })?; + let mut response = client.get(url.clone()).send().await.map_err(|error| { + UseError::new( + "use.extension.registry_download_failed", + format!("Failed to download the bootstrap TUF root: {error}"), + ) + })?; + validate_download_url(response.url())?; + if !response.status().is_success() { + return Err(UseError::new( + "use.extension.registry_download_failed", + format!( + "Bootstrap TUF root download returned HTTP {}.", + response.status() + ), + )); + } + if response + .content_length() + .is_some_and(|length| length > MAX_BOOTSTRAP_ROOT_BYTES) + { + return Err(UseError::new( + "use.extension.registry_root_invalid", + "The bootstrap TUF root exceeds the one MiB limit.", + )); + } + let mut bytes = Vec::with_capacity( + response + .content_length() + .unwrap_or_default() + .min(MAX_BOOTSTRAP_ROOT_BYTES) as usize, + ); + while let Some(chunk) = response.chunk().await.map_err(|error| { + UseError::new( + "use.extension.registry_download_failed", + format!("Failed to read the bootstrap TUF root: {error}"), + ) + })? { + if bytes.len().saturating_add(chunk.len()) as u64 > MAX_BOOTSTRAP_ROOT_BYTES { + return Err(UseError::new( + "use.extension.registry_root_invalid", + "The bootstrap TUF root exceeds the one MiB limit.", + )); + } + bytes.extend_from_slice(&chunk); + } + Ok(bytes) +} + +fn verify_root_digest(registry: &TrustedRegistry, bytes: &[u8]) -> UseResult<()> { + let actual = format!("{:x}", Sha256::digest(bytes)); + if actual == registry.root_sha256 { + return Ok(()); + } + Err(UseError::new( + "use.extension.registry_root_mismatch", + format!( + "Registry '{}' bootstrap root does not match its pinned SHA-256.", + registry.name + ), + ) + .with_detail("expected", registry.root_sha256.clone()) + .with_detail("actual", actual)) +} + +async fn write_bootstrap_root(path: &Path, bytes: &[u8]) -> UseResult<()> { + let parent = path.parent().ok_or_else(|| { + UseError::new( + "use.extension.registry_path_invalid", + "The bootstrap TUF root cache has no parent directory.", + ) + })?; + let temporary = parent.join(format!(".root-{}.tmp", unique_suffix())); + let mut options = fs::OpenOptions::new(); + options.create_new(true).write(true); + let mut file = options + .open(&temporary) + .await + .map_err(|error| io_error("create bootstrap TUF root cache", &temporary, error))?; + if let Err(error) = file.write_all(bytes).await { + let _ = fs::remove_file(&temporary).await; + return Err(io_error( + "write bootstrap TUF root cache", + &temporary, + error, + )); + } + if let Err(error) = file.sync_all().await { + let _ = fs::remove_file(&temporary).await; + return Err(io_error("sync bootstrap TUF root cache", &temporary, error)); + } + drop(file); + if let Err(error) = activate_temporary_file( + temporary.clone(), + path.to_path_buf(), + "activate bootstrap TUF root cache", + ) + .await + { + let _ = fs::remove_file(&temporary).await; + return Err(error); + } + sync_parent_directory(parent, "TUF metadata").await +} + +fn normalize_registry_url(value: &str) -> UseResult { + let mut url = Url::parse(value).map_err(|error| { + UseError::new( + "use.extension.registry_url_invalid", + format!("Invalid registry URL: {error}"), + ) + })?; + validate_download_url(&url)?; + if !url.username().is_empty() + || url.password().is_some() + || url.query().is_some() + || url.fragment().is_some() + { + return Err(UseError::new( + "use.extension.registry_url_invalid", + "Registry URLs must not contain credentials, query parameters, or fragments.", + )); + } + if !url.path().ends_with('/') { + let path = format!("{}/", url.path()); + url.set_path(&path); + } + Ok(url) +} + +fn validate_download_url(url: &Url) -> UseResult<()> { + let https = url.scheme() == "https"; + let loopback_http = url.scheme() == "http" + && url.host_str().is_some_and(|host| { + host.eq_ignore_ascii_case("localhost") + || host + .parse::() + .is_ok_and(|ip| ip.is_loopback()) + }); + if https || loopback_http { + Ok(()) + } else { + Err(UseError::new( + "use.extension.registry_url_invalid", + "Registry downloads require HTTPS; HTTP is accepted only on loopback for local testing.", + )) + } +} + +fn validate_registry_name(name: &str) -> UseResult<()> { + let mut characters = name.chars(); + if characters + .next() + .is_some_and(|character| character.is_ascii_lowercase()) + && characters.all(|character| { + character.is_ascii_lowercase() || character.is_ascii_digit() || character == '-' + }) + { + Ok(()) + } else { + Err(UseError::new( + "use.extension.registry_name_invalid", + "Registry names use lowercase letters, digits, and hyphens and start with a letter.", + )) + } +} + +fn normalize_sha256(value: &str, label: &str) -> UseResult { + let value = value.strip_prefix("sha256:").unwrap_or(value); + if value.len() == 64 + && value + .bytes() + .all(|byte| byte.is_ascii_hexdigit() && !byte.is_ascii_uppercase()) + { + Ok(value.to_string()) + } else { + Err(UseError::new( + "use.extension.registry_digest_invalid", + format!("The {label} must be exactly 64 lowercase hexadecimal characters."), + )) + } +} + +fn hex_lower(bytes: &[u8]) -> String { + let mut output = String::with_capacity(bytes.len() * 2); + for byte in bytes { + use std::fmt::Write as _; + let _ = write!(output, "{byte:02x}"); + } + output +} + +#[cfg(test)] +#[path = "tuf_test_support.rs"] +mod test_support; + +#[cfg(test)] +#[path = "remote_tests.rs"] +mod tests; diff --git a/crates/extension/src/remote_tests.rs b/crates/extension/src/remote_tests.rs new file mode 100644 index 00000000..0b538dc8 --- /dev/null +++ b/crates/extension/src/remote_tests.rs @@ -0,0 +1,310 @@ +use std::path::PathBuf; + +use super::test_support::{ + extension_archive, find_subslice, TestRepository, TestServer, EXPIRED, FUTURE, PACKAGE_VERSION, +}; +use super::*; +use crate::{ExtensionPaths, ExtensionRegistry, ExtensionTrust}; + +#[tokio::test] +async fn tuf_refresh_verifies_metadata_without_downloading_targets() { + let repository = TestRepository::new(extension_archive(PACKAGE_VERSION), 7, FUTURE); + let server = TestServer::start(repository.routes.clone()); + let temp = tempfile::tempdir().unwrap(); + let trusted = trusted_registry(&server, &repository, temp.path().join("tuf")); + + let metadata = refresh_remote_registry(&trusted).await.unwrap(); + + assert_eq!(metadata.registry_name, "fixture"); + assert_eq!(metadata.root_version, 1); + assert_eq!(metadata.timestamp_version, 7); + assert_eq!(metadata.snapshot_version, 7); + assert_eq!(metadata.targets_version, 7); + assert_eq!(metadata.package_targets, 1); + assert!(server + .requests() + .iter() + .all(|request| !request.starts_with("/targets/"))); +} + +#[tokio::test] +async fn tuf_install_records_signed_provenance_and_converges() { + let archive = extension_archive(PACKAGE_VERSION); + let repository = TestRepository::new(archive, 1, FUTURE); + let server = TestServer::start(repository.routes.clone()); + let temp = tempfile::tempdir().unwrap(); + let trusted = trusted_registry(&server, &repository, temp.path().join("tuf")); + + let prepared = prepare_remote_package(&trusted, "acme/slack", None, "stable", None) + .await + .unwrap(); + let digest = prepared.resolved().plan_digest().unwrap(); + drop(prepared); + assert!(server + .requests() + .iter() + .all(|request| !request.starts_with("/targets/"))); + + let paths = ExtensionPaths::new( + temp.path().join("data"), + temp.path().join("extension-state"), + ); + let registry = ExtensionRegistry::new(paths); + let installed = registry + .install_remote("acme/slack", &trusted, None, "stable", Some(&digest), false) + .await + .unwrap(); + assert!(installed.changed); + assert_eq!( + installed.extension.receipt.trust, + ExtensionTrust::RegistryTuf + ); + let provenance = installed.extension.receipt.registry.as_ref().unwrap(); + assert_eq!(provenance.package_id, "acme/slack"); + assert_eq!(provenance.version, PACKAGE_VERSION); + assert_eq!(provenance.sha256, repository.target_sha256); + assert!(installed.extension.cli_executable().unwrap().is_file()); + + server.clear_requests(); + let second = registry + .install_remote("acme/slack", &trusted, None, "stable", Some(&digest), false) + .await + .unwrap(); + assert!(!second.changed); + assert_eq!(registry.list().await.unwrap().len(), 1); + assert!(server + .requests() + .iter() + .all(|request| !request.starts_with("/targets/"))); +} + +#[tokio::test] +async fn tuf_convergence_refreshes_signed_provenance_without_downloading_the_target() { + let archive = extension_archive(PACKAGE_VERSION); + let first_repository = TestRepository::new(archive.clone(), 1, FUTURE); + let server = TestServer::start(first_repository.routes.clone()); + let temp = tempfile::tempdir().unwrap(); + let trusted = trusted_registry(&server, &first_repository, temp.path().join("tuf")); + let paths = ExtensionPaths::new( + temp.path().join("data"), + temp.path().join("extension-state"), + ); + let registry = ExtensionRegistry::new(paths); + registry + .install_remote("acme/slack", &trusted, None, "stable", None, false) + .await + .unwrap(); + + let second_repository = TestRepository::new(archive, 2, FUTURE); + assert_eq!( + second_repository.target_sha256, + first_repository.target_sha256 + ); + server.replace_routes(second_repository.routes); + server.clear_requests(); + + let converged = registry + .install_remote("acme/slack", &trusted, None, "stable", None, false) + .await + .unwrap(); + + assert!(!converged.changed); + let provenance = converged.extension.receipt.registry.unwrap(); + assert_eq!(provenance.timestamp_version, 2); + assert_eq!(provenance.snapshot_version, 2); + assert_eq!(provenance.targets_version, 2); + assert!(server + .requests() + .iter() + .all(|request| !request.starts_with("/targets/"))); +} + +#[tokio::test] +async fn tuf_install_rejects_modified_installed_content_before_dispatch_or_convergence() { + let repository = TestRepository::new(extension_archive(PACKAGE_VERSION), 1, FUTURE); + let server = TestServer::start(repository.routes.clone()); + let temp = tempfile::tempdir().unwrap(); + let trusted = trusted_registry(&server, &repository, temp.path().join("tuf")); + let paths = ExtensionPaths::new( + temp.path().join("data"), + temp.path().join("extension-state"), + ); + let registry = ExtensionRegistry::new(paths); + let installed = registry + .install_remote("acme/slack", &trusted, None, "stable", None, false) + .await + .unwrap(); + std::fs::write( + installed.extension.cli_executable().unwrap(), + b"modified executable", + ) + .unwrap(); + + let dispatch_error = match registry.acquire_route("slack").await { + Err(error) => error, + Ok(_) => panic!("modified signed content must not be dispatched"), + }; + assert_eq!(dispatch_error.code, "use.extension.package_digest_mismatch"); + + server.clear_requests(); + let convergence_error = registry + .install_remote("acme/slack", &trusted, None, "stable", None, false) + .await + .unwrap_err(); + assert_eq!( + convergence_error.code, + "use.extension.package_digest_mismatch" + ); + assert!(server + .requests() + .iter() + .all(|request| !request.starts_with("/targets/"))); +} + +#[tokio::test] +async fn tuf_receipt_requires_an_expanded_package_digest() { + let repository = TestRepository::new(extension_archive(PACKAGE_VERSION), 1, FUTURE); + let server = TestServer::start(repository.routes.clone()); + let temp = tempfile::tempdir().unwrap(); + let trusted = trusted_registry(&server, &repository, temp.path().join("tuf")); + let paths = ExtensionPaths::new( + temp.path().join("data"), + temp.path().join("extension-state"), + ); + let registry = ExtensionRegistry::new(paths); + registry + .install_remote("acme/slack", &trusted, None, "stable", None, false) + .await + .unwrap(); + + let receipt_path = registry.paths().receipt_path("acme/slack"); + let mut receipt: serde_json::Value = + serde_json::from_slice(&std::fs::read(&receipt_path).unwrap()).unwrap(); + receipt.as_object_mut().unwrap().remove("packageSha256"); + std::fs::write(&receipt_path, serde_json::to_vec_pretty(&receipt).unwrap()).unwrap(); + + let error = registry.get("acme/slack").await.unwrap_err(); + assert_eq!(error.code, "use.extension.receipt_invalid"); +} + +#[tokio::test] +async fn reviewed_registry_plan_fails_before_target_download() { + let repository = TestRepository::new(extension_archive(PACKAGE_VERSION), 1, FUTURE); + let server = TestServer::start(repository.routes.clone()); + let temp = tempfile::tempdir().unwrap(); + let trusted = trusted_registry(&server, &repository, temp.path().join("tuf")); + + let error = prepare_remote_package( + &trusted, + "acme/slack", + None, + "stable", + Some(&"0".repeat(64)), + ) + .await + .unwrap_err(); + + assert_eq!(error.code, "use.extension.registry_plan_mismatch"); + assert!(server + .requests() + .iter() + .all(|request| !request.starts_with("/targets/"))); +} + +#[tokio::test] +async fn tuf_rejects_wrong_root_and_tampered_target() { + let archive = extension_archive(PACKAGE_VERSION); + let repository = TestRepository::new(archive, 1, FUTURE); + let server = TestServer::start(repository.routes.clone()); + let temp = tempfile::tempdir().unwrap(); + let wrong = TrustedRegistry::new( + "fixture", + server.base_url(), + "f".repeat(64), + None, + temp.path().join("wrong-root"), + ) + .unwrap(); + let error = prepare_remote_package(&wrong, "acme/slack", None, "stable", None) + .await + .unwrap_err(); + assert_eq!(error.code, "use.extension.registry_root_mismatch"); + + let mut routes = repository.routes.clone(); + routes.insert( + format!("/targets/{}", repository.target_name), + b"tampered archive".to_vec(), + ); + let tampered_server = TestServer::start(routes); + let trusted = trusted_registry( + &tampered_server, + &repository, + temp.path().join("tampered-target"), + ); + let prepared = prepare_remote_package(&trusted, "acme/slack", None, "stable", None) + .await + .unwrap(); + let error = prepared.download().await.unwrap_err(); + assert_eq!(error.code, "use.extension.registry_download_failed"); +} + +#[tokio::test] +async fn tuf_rejects_metadata_tampering_expiration_and_rollback() { + let archive = extension_archive(PACKAGE_VERSION); + let version_two = TestRepository::new(archive.clone(), 2, FUTURE); + let server_two = TestServer::start(version_two.routes.clone()); + let temp = tempfile::tempdir().unwrap(); + let datastore = temp.path().join("rollback-state"); + let trusted_two = trusted_registry(&server_two, &version_two, datastore.clone()); + prepare_remote_package(&trusted_two, "acme/slack", None, "stable", None) + .await + .unwrap(); + + let version_one = TestRepository::new(archive.clone(), 1, FUTURE); + assert_eq!(version_one.root_sha256, version_two.root_sha256); + let server_one = TestServer::start(version_one.routes.clone()); + let trusted_one = trusted_registry(&server_one, &version_one, datastore); + let rollback = prepare_remote_package(&trusted_one, "acme/slack", None, "stable", None) + .await + .unwrap_err(); + assert_eq!(rollback.code, "use.extension.registry_untrusted"); + + let expired = TestRepository::new(archive.clone(), 1, EXPIRED); + let expired_server = TestServer::start(expired.routes.clone()); + let expired_registry = + trusted_registry(&expired_server, &expired, temp.path().join("expired-state")); + let error = prepare_remote_package(&expired_registry, "acme/slack", None, "stable", None) + .await + .unwrap_err(); + assert_eq!(error.code, "use.extension.registry_untrusted"); + + let mut tampered_routes = version_one.routes.clone(); + let targets = tampered_routes.get_mut("/metadata/targets.json").unwrap(); + let position = find_subslice(targets, b"stable").unwrap(); + targets[position..position + 6].copy_from_slice(b"nightl"); + let tampered_server = TestServer::start(tampered_routes); + let tampered_registry = trusted_registry( + &tampered_server, + &version_one, + temp.path().join("tampered-metadata"), + ); + let error = prepare_remote_package(&tampered_registry, "acme/slack", None, "stable", None) + .await + .unwrap_err(); + assert_eq!(error.code, "use.extension.registry_untrusted"); +} + +fn trusted_registry( + server: &TestServer, + repository: &TestRepository, + datastore: PathBuf, +) -> TrustedRegistry { + TrustedRegistry::new( + "fixture", + server.base_url(), + &repository.root_sha256, + None, + datastore, + ) + .unwrap() +} diff --git a/crates/extension/src/source.rs b/crates/extension/src/source.rs new file mode 100644 index 00000000..e8132daf --- /dev/null +++ b/crates/extension/src/source.rs @@ -0,0 +1,708 @@ +use std::collections::BTreeSet; +use std::fs::{File, OpenOptions}; +use std::io::{self, Read, Write}; +use std::path::{Component, Path, PathBuf}; + +use a3s_use_core::{UseError, UseResult}; +use tempfile::TempDir; +use tokio::fs; + +use super::package::{io_error, MANIFEST_NAME, MAX_PACKAGE_BYTES, MAX_PACKAGE_FILES}; + +const MAX_ARCHIVE_BYTES: u64 = 512 * 1024 * 1024; +const MAX_PATH_BYTES: usize = 4_096; +const MAX_PATH_DEPTH: usize = 32; + +#[derive(Clone, Copy)] +enum ArchiveKind { + TarGz, + Zip, +} + +struct ExtractedEntry { + relative: PathBuf, + file: bool, +} + +/// One validated local package source kept alive through installation. +#[derive(Debug)] +pub(crate) struct PreparedPackageSource { + root: PathBuf, + _temporary: Option, +} + +impl PreparedPackageSource { + pub(crate) fn root(&self) -> &Path { + &self.root + } +} + +pub(crate) async fn prepare_package_source(source: &Path) -> UseResult { + let source = fs::canonicalize(source) + .await + .map_err(|error| io_error("resolve extension package", source, error))?; + let metadata = fs::metadata(&source) + .await + .map_err(|error| io_error("inspect extension package", &source, error))?; + if metadata.is_dir() { + return Ok(PreparedPackageSource { + root: source, + _temporary: None, + }); + } + if !metadata.is_file() { + return Err(UseError::new( + "use.extension.package_unsupported", + "The local extension source must be a package directory, .tar.gz, .tgz, or .zip archive.", + )); + } + if metadata.len() > MAX_ARCHIVE_BYTES { + return Err(UseError::new( + "use.extension.package_too_large", + format!( + "The extension archive exceeds the {MAX_ARCHIVE_BYTES} byte compressed-size limit." + ), + )); + } + let kind = archive_kind(&source)?; + let temporary = tokio::task::spawn_blocking(tempfile::tempdir) + .await + .map_err(|error| { + UseError::new( + "use.extension.io", + format!("Failed to create extension archive staging task: {error}"), + ) + })? + .map_err(|error| io_error("create extension archive staging directory", &source, error))?; + let extraction_root = temporary.path().join("package"); + let blocking_source = source.clone(); + let blocking_root = extraction_root.clone(); + let package_relative = tokio::task::spawn_blocking(move || { + extract_archive(&blocking_source, &blocking_root, kind) + }) + .await + .map_err(|error| { + UseError::new( + "use.extension.package_archive_invalid", + format!("Extension archive extraction task failed: {error}"), + ) + })??; + let extraction_root = fs::canonicalize(&extraction_root).await.map_err(|error| { + io_error( + "resolve extension archive staging directory", + &extraction_root, + error, + ) + })?; + let root = extraction_root.join(package_relative); + let root = fs::canonicalize(&root) + .await + .map_err(|error| io_error("resolve extracted extension package", &root, error))?; + if !root.starts_with(&extraction_root) { + return Err(UseError::new( + "use.extension.path_escape", + "The extracted extension package root escapes its staging directory.", + )); + } + Ok(PreparedPackageSource { + root, + _temporary: Some(temporary), + }) +} + +fn archive_kind(path: &Path) -> UseResult { + let name = path + .file_name() + .and_then(|value| value.to_str()) + .unwrap_or_default() + .to_ascii_lowercase(); + if name.ends_with(".tar.gz") || name.ends_with(".tgz") { + Ok(ArchiveKind::TarGz) + } else if name.ends_with(".zip") { + Ok(ArchiveKind::Zip) + } else { + Err(UseError::new( + "use.extension.package_unsupported", + "Extension package archives must use .tar.gz, .tgz, or .zip.", + )) + } +} + +fn extract_archive(source: &Path, target: &Path, kind: ArchiveKind) -> UseResult { + std::fs::create_dir_all(target) + .map_err(|error| archive_io("create extraction directory", target, error))?; + let entries = match kind { + ArchiveKind::TarGz => extract_tar_gz(source, target)?, + ArchiveKind::Zip => extract_zip(source, target)?, + }; + resolve_package_root(&entries) +} + +fn extract_tar_gz(source: &Path, target: &Path) -> UseResult> { + let file = File::open(source).map_err(|error| archive_io("open", source, error))?; + let decoder = flate2::read::GzDecoder::new(file); + let mut archive = tar::Archive::new(decoder); + let mut extracted = Vec::new(); + let mut seen = BTreeSet::new(); + let mut extracted_bytes = 0_u64; + let entries = archive + .entries() + .map_err(|error| archive_invalid(format!("Failed to read tar entries: {error}")))?; + for (entry_count, entry) in entries.enumerate() { + if entry_count >= MAX_PACKAGE_FILES { + return Err(package_limit_error()); + } + let mut entry = + entry.map_err(|error| archive_invalid(format!("Failed to read tar entry: {error}")))?; + let entry_path = entry + .path() + .map_err(|error| archive_invalid(format!("Failed to read tar entry path: {error}")))? + .into_owned(); + let entry_type = entry.header().entry_type(); + if ignored_macos_metadata_path(&entry_path)? { + if entry_type.is_dir() { + continue; + } + if entry_type.is_file() { + let remaining = MAX_PACKAGE_BYTES.saturating_sub(extracted_bytes); + extracted_bytes = extracted_bytes.saturating_add(copy_bounded( + &mut entry, + &mut io::sink(), + remaining, + &entry_path, + )?); + continue; + } + } + let Some(relative) = sanitized_relative_path(&entry_path)? else { + if entry_type.is_dir() { + continue; + } + return Err(archive_invalid( + "The archive contains a non-directory root entry.", + )); + }; + if !seen.insert(relative.clone()) { + return Err(archive_invalid(format!( + "The archive contains duplicate entry '{}'.", + relative.display() + ))); + } + let output = target.join(&relative); + if entry_type.is_dir() { + std::fs::create_dir_all(&output) + .map_err(|error| archive_io("create archive directory", &output, error))?; + extracted.push(ExtractedEntry { + relative, + file: false, + }); + } else if entry_type.is_file() { + let remaining = MAX_PACKAGE_BYTES.saturating_sub(extracted_bytes); + if entry.size() > remaining { + return Err(package_limit_error()); + } + if let Some(parent) = output.parent() { + std::fs::create_dir_all(parent) + .map_err(|error| archive_io("create archive parent", parent, error))?; + } + let mut output_file = OpenOptions::new() + .create_new(true) + .write(true) + .open(&output) + .map_err(|error| archive_io("create archive file", &output, error))?; + extracted_bytes = extracted_bytes.saturating_add(copy_bounded( + &mut entry, + &mut output_file, + remaining, + &output, + )?); + apply_unix_mode(&output, entry.header().mode().ok())?; + extracted.push(ExtractedEntry { + relative, + file: true, + }); + } else if entry_type.is_symlink() || entry_type.is_hard_link() { + return Err(UseError::new( + "use.extension.package_symlink", + format!( + "Extension archive entry '{}' is a link.", + relative.display() + ), + )); + } else { + return Err(UseError::new( + "use.extension.package_entry_invalid", + format!( + "Extension archive entry '{}' is not a regular file or directory.", + relative.display() + ), + )); + } + } + Ok(extracted) +} + +fn extract_zip(source: &Path, target: &Path) -> UseResult> { + let file = File::open(source).map_err(|error| archive_io("open", source, error))?; + let mut archive = zip::ZipArchive::new(file) + .map_err(|error| archive_invalid(format!("Failed to read ZIP archive: {error}")))?; + if archive.len() > MAX_PACKAGE_FILES { + return Err(package_limit_error()); + } + let mut extracted = Vec::new(); + let mut seen = BTreeSet::new(); + let mut extracted_bytes = 0_u64; + for index in 0..archive.len() { + let mut entry = archive.by_index(index).map_err(|error| { + archive_invalid(format!("Failed to read ZIP entry {index}: {error}")) + })?; + if entry.is_symlink() { + return Err(UseError::new( + "use.extension.package_symlink", + format!("Extension archive entry '{}' is a link.", entry.name()), + )); + } + let enclosed = entry.enclosed_name().ok_or_else(|| { + UseError::new( + "use.extension.path_escape", + format!( + "Extension archive entry '{}' escapes the package.", + entry.name() + ), + ) + })?; + if ignored_macos_metadata_path(&enclosed)? { + if entry.is_dir() { + continue; + } + if entry.is_file() { + let remaining = MAX_PACKAGE_BYTES.saturating_sub(extracted_bytes); + extracted_bytes = extracted_bytes.saturating_add(copy_bounded( + &mut entry, + &mut io::sink(), + remaining, + &enclosed, + )?); + continue; + } + } + let Some(relative) = sanitized_relative_path(&enclosed)? else { + if entry.is_dir() { + continue; + } + return Err(archive_invalid( + "The ZIP archive contains a non-directory root entry.", + )); + }; + if !seen.insert(relative.clone()) { + return Err(archive_invalid(format!( + "The ZIP archive contains duplicate entry '{}'.", + relative.display() + ))); + } + let output = target.join(&relative); + if entry.is_dir() { + std::fs::create_dir_all(&output) + .map_err(|error| archive_io("create ZIP directory", &output, error))?; + extracted.push(ExtractedEntry { + relative, + file: false, + }); + } else if entry.is_file() { + let remaining = MAX_PACKAGE_BYTES.saturating_sub(extracted_bytes); + if entry.size() > remaining { + return Err(package_limit_error()); + } + if let Some(parent) = output.parent() { + std::fs::create_dir_all(parent) + .map_err(|error| archive_io("create ZIP parent", parent, error))?; + } + let mut output_file = OpenOptions::new() + .create_new(true) + .write(true) + .open(&output) + .map_err(|error| archive_io("create ZIP file", &output, error))?; + extracted_bytes = extracted_bytes.saturating_add(copy_bounded( + &mut entry, + &mut output_file, + remaining, + &output, + )?); + apply_unix_mode(&output, entry.unix_mode())?; + extracted.push(ExtractedEntry { + relative, + file: true, + }); + } else { + return Err(UseError::new( + "use.extension.package_entry_invalid", + format!("Extension ZIP entry '{}' is unsupported.", entry.name()), + )); + } + } + Ok(extracted) +} + +fn ignored_macos_metadata_path(path: &Path) -> UseResult { + if path.as_os_str().is_empty() { + return Err(archive_invalid("The archive contains an empty entry path.")); + } + let encoded = path.to_str().ok_or_else(|| { + archive_invalid("Extension archive paths must be valid UTF-8 for portability.") + })?; + if encoded.len() > MAX_PATH_BYTES { + return Err(archive_invalid(format!( + "Extension archive path '{}' is not portable.", + path.display() + ))); + } + + let mut first = None; + let mut last = None; + let mut depth = 0_usize; + for component in path.components() { + match component { + Component::Normal(segment) => { + depth += 1; + if depth > MAX_PATH_DEPTH { + return Err(archive_invalid(format!( + "Extension archive path '{}' exceeds the depth limit.", + path.display() + ))); + } + let segment = segment.to_str().ok_or_else(|| { + archive_invalid("Extension archive paths must be valid UTF-8 for portability.") + })?; + first.get_or_insert(segment); + last = Some(segment); + } + Component::CurDir => {} + Component::ParentDir | Component::RootDir | Component::Prefix(_) => { + return Err(UseError::new( + "use.extension.path_escape", + format!( + "Extension archive path '{}' escapes the package.", + path.display() + ), + )); + } + } + } + + Ok(first == Some("__MACOSX") || last.is_some_and(|segment| segment.starts_with("._"))) +} + +pub(crate) fn sanitized_relative_path(path: &Path) -> UseResult> { + if path.as_os_str().is_empty() { + return Err(archive_invalid("The archive contains an empty entry path.")); + } + let encoded = path.to_str().ok_or_else(|| { + archive_invalid("Extension archive paths must be valid UTF-8 for portability.") + })?; + if encoded.len() > MAX_PATH_BYTES { + return Err(archive_invalid(format!( + "Extension archive path '{}' is not portable.", + path.display() + ))); + } + let mut sanitized = PathBuf::new(); + let mut depth = 0_usize; + for component in path.components() { + match component { + Component::Normal(segment) => { + depth += 1; + if depth > MAX_PATH_DEPTH { + return Err(archive_invalid(format!( + "Extension archive path '{}' exceeds the depth limit.", + path.display() + ))); + } + validate_portable_segment(segment, path)?; + sanitized.push(segment); + } + Component::CurDir => {} + Component::ParentDir | Component::RootDir | Component::Prefix(_) => { + return Err(UseError::new( + "use.extension.path_escape", + format!( + "Extension archive path '{}' escapes the package.", + path.display() + ), + )); + } + } + } + if sanitized.as_os_str().is_empty() { + Ok(None) + } else { + Ok(Some(sanitized)) + } +} + +fn validate_portable_segment(segment: &std::ffi::OsStr, path: &Path) -> UseResult<()> { + let segment = segment.to_str().ok_or_else(|| { + archive_invalid("Extension archive paths must be valid UTF-8 for portability.") + })?; + if segment.ends_with(['.', ' ']) + || segment + .chars() + .any(|character| character.is_control() || r#"<>:"/\|?*"#.contains(character)) + { + return Err(archive_invalid(format!( + "Extension archive path '{}' is not portable.", + path.display() + ))); + } + let device = segment + .split('.') + .next() + .unwrap_or_default() + .to_ascii_uppercase(); + let reserved = matches!(device.as_str(), "CON" | "PRN" | "AUX" | "NUL") + || device + .strip_prefix("COM") + .or_else(|| device.strip_prefix("LPT")) + .is_some_and(|number| { + matches!(number, "1" | "2" | "3" | "4" | "5" | "6" | "7" | "8" | "9") + }); + if reserved { + return Err(archive_invalid(format!( + "Extension archive path '{}' uses a reserved device name.", + path.display() + ))); + } + Ok(()) +} + +fn copy_bounded( + reader: &mut impl Read, + writer: &mut impl Write, + remaining: u64, + path: &Path, +) -> UseResult { + let mut bounded = reader.take(remaining.saturating_add(1)); + let copied = io::copy(&mut bounded, writer) + .map_err(|error| archive_io("extract archive file", path, error))?; + if copied > remaining { + return Err(package_limit_error()); + } + Ok(copied) +} + +fn resolve_package_root(entries: &[ExtractedEntry]) -> UseResult { + let manifests = entries + .iter() + .filter(|entry| { + entry.file + && entry + .relative + .file_name() + .is_some_and(|name| name == MANIFEST_NAME) + }) + .collect::>(); + let [manifest] = manifests.as_slice() else { + return Err(UseError::new( + "use.extension.package_layout_invalid", + format!("Extension archives must contain exactly one regular {MANIFEST_NAME} file."), + )); + }; + let root = manifest + .relative + .parent() + .map(Path::to_path_buf) + .unwrap_or_default(); + if !root.as_os_str().is_empty() + && entries + .iter() + .any(|entry| !entry.relative.starts_with(&root)) + { + return Err(UseError::new( + "use.extension.package_layout_invalid", + "Extension archive entries must all belong to the directory containing its manifest.", + )); + } + Ok(root) +} + +#[cfg(unix)] +fn apply_unix_mode(path: &Path, mode: Option) -> UseResult<()> { + use std::os::unix::fs::PermissionsExt; + + if let Some(mode) = mode { + std::fs::set_permissions(path, std::fs::Permissions::from_mode(mode & 0o777)) + .map_err(|error| archive_io("set archive file permissions", path, error))?; + } + Ok(()) +} + +#[cfg(not(unix))] +fn apply_unix_mode(_path: &Path, _mode: Option) -> UseResult<()> { + Ok(()) +} + +fn archive_io(action: &str, path: &Path, error: io::Error) -> UseError { + UseError::new( + "use.extension.package_archive_invalid", + format!( + "Failed to {action} extension archive entry '{}': {error}", + path.display() + ), + ) +} + +fn archive_invalid(message: impl Into) -> UseError { + UseError::new("use.extension.package_archive_invalid", message) +} + +fn package_limit_error() -> UseError { + UseError::new( + "use.extension.package_too_large", + "The extension package exceeds the local installation limits.", + ) +} + +#[cfg(test)] +mod tests { + use std::io::Write; + + use super::*; + + #[tokio::test] + async fn tar_package_accepts_an_explicit_current_directory_root() { + let temp = tempfile::tempdir().unwrap(); + let archive_path = temp.path().join("package.tar.gz"); + { + let file = File::create(&archive_path).unwrap(); + let encoder = flate2::write::GzEncoder::new(file, flate2::Compression::default()); + let mut builder = tar::Builder::new(encoder); + + let mut root = tar::Header::new_gnu(); + root.set_path(".").unwrap(); + root.set_entry_type(tar::EntryType::Directory); + root.set_size(0); + root.set_mode(0o755); + root.set_cksum(); + builder.append(&root, io::empty()).unwrap(); + + let manifest = b"extension fixture"; + let mut header = tar::Header::new_gnu(); + header.set_path(format!("./{MANIFEST_NAME}")).unwrap(); + header.set_size(manifest.len() as u64); + header.set_mode(0o644); + header.set_cksum(); + builder.append(&header, &manifest[..]).unwrap(); + builder.finish().unwrap(); + } + + let prepared = prepare_package_source(&archive_path).await.unwrap(); + assert_eq!( + std::fs::read(prepared.root().join(MANIFEST_NAME)).unwrap(), + b"extension fixture" + ); + } + + #[tokio::test] + async fn tar_package_ignores_bounded_macos_appledouble_metadata() { + let temp = tempfile::tempdir().unwrap(); + let archive_path = temp.path().join("package.tar.gz"); + { + let file = File::create(&archive_path).unwrap(); + let encoder = flate2::write::GzEncoder::new(file, flate2::Compression::default()); + let mut builder = tar::Builder::new(encoder); + + let metadata = b"appledouble"; + let mut metadata_header = tar::Header::new_gnu(); + metadata_header.set_path("./._.").unwrap(); + metadata_header.set_size(metadata.len() as u64); + metadata_header.set_mode(0o644); + metadata_header.set_cksum(); + builder.append(&metadata_header, &metadata[..]).unwrap(); + + let manifest = b"extension fixture"; + let mut manifest_header = tar::Header::new_gnu(); + manifest_header + .set_path(format!("./{MANIFEST_NAME}")) + .unwrap(); + manifest_header.set_size(manifest.len() as u64); + manifest_header.set_mode(0o644); + manifest_header.set_cksum(); + builder.append(&manifest_header, &manifest[..]).unwrap(); + builder.finish().unwrap(); + } + + let prepared = prepare_package_source(&archive_path).await.unwrap(); + assert_eq!( + std::fs::read(prepared.root().join(MANIFEST_NAME)).unwrap(), + b"extension fixture" + ); + assert!(!prepared.root().join("._.").exists()); + } + + #[tokio::test] + async fn zip_package_rejects_parent_traversal() { + let temp = tempfile::tempdir().unwrap(); + let archive_path = temp.path().join("escape.zip"); + { + let file = File::create(&archive_path).unwrap(); + let mut writer = zip::ZipWriter::new(file); + writer + .start_file( + "../a3s-use-extension.acl", + zip::write::SimpleFileOptions::default(), + ) + .unwrap(); + writer.write_all(b"escape").unwrap(); + writer.finish().unwrap(); + } + + let error = prepare_package_source(&archive_path).await.unwrap_err(); + assert_eq!(error.code, "use.extension.path_escape"); + } + + #[tokio::test] + async fn tar_package_rejects_symbolic_links() { + let temp = tempfile::tempdir().unwrap(); + let archive_path = temp.path().join("link.tar.gz"); + { + let file = File::create(&archive_path).unwrap(); + let encoder = flate2::write::GzEncoder::new(file, flate2::Compression::default()); + let mut builder = tar::Builder::new(encoder); + + let manifest = b"extension fixture"; + let mut manifest_header = tar::Header::new_gnu(); + manifest_header + .set_path(format!("package/{MANIFEST_NAME}")) + .unwrap(); + manifest_header.set_size(manifest.len() as u64); + manifest_header.set_mode(0o644); + manifest_header.set_cksum(); + builder.append(&manifest_header, &manifest[..]).unwrap(); + + let mut link = tar::Header::new_gnu(); + link.set_entry_type(tar::EntryType::Symlink); + link.set_path("package/escape").unwrap(); + link.set_link_name("../../outside").unwrap(); + link.set_size(0); + link.set_cksum(); + builder.append(&link, io::empty()).unwrap(); + builder.finish().unwrap(); + } + + let error = prepare_package_source(&archive_path).await.unwrap_err(); + assert_eq!(error.code, "use.extension.package_symlink"); + } + + #[test] + fn archive_paths_reject_cross_platform_escapes_and_device_names() { + for path in ["C:/escape", "..\\escape", "package/CON", "package/name. "] { + assert!( + sanitized_relative_path(Path::new(path)).is_err(), + "accepted unsafe path {path}" + ); + } + assert_eq!( + sanitized_relative_path(Path::new("./package/bin/tool")).unwrap(), + Some(PathBuf::from("package/bin/tool")) + ); + } +} diff --git a/crates/extension/src/tuf_test_support.rs b/crates/extension/src/tuf_test_support.rs new file mode 100644 index 00000000..06767097 --- /dev/null +++ b/crates/extension/src/tuf_test_support.rs @@ -0,0 +1,324 @@ +#![allow(dead_code)] + +use std::collections::HashMap; +use std::io::{Read, Write}; +use std::net::{Shutdown, TcpListener, TcpStream}; +use std::sync::atomic::{AtomicBool, Ordering}; +use std::sync::{Arc, Mutex}; +use std::thread::JoinHandle; +use std::time::Duration; + +use olpc_cjson::CanonicalFormatter; +use ring::signature::{Ed25519KeyPair, KeyPair}; +use serde::Serialize; +use serde_json::{json, Map, Value}; +use sha2::{Digest, Sha256}; + +pub(crate) const FUTURE: &str = "2999-01-01T00:00:00Z"; +pub(crate) const EXPIRED: &str = "2000-01-01T00:00:00Z"; +pub(crate) const PACKAGE_VERSION: &str = "0.1.1"; + +pub(crate) struct TestRepository { + pub(crate) routes: HashMap>, + pub(crate) root_sha256: String, + pub(crate) target_name: String, + pub(crate) target_sha256: String, +} + +impl TestRepository { + pub(crate) fn new(archive: Vec, metadata_version: u64, expires: &str) -> Self { + Self::with_package_version(archive, PACKAGE_VERSION, metadata_version, expires) + } + + pub(crate) fn with_package_version( + archive: Vec, + package_version: &str, + metadata_version: u64, + expires: &str, + ) -> Self { + let key = Ed25519KeyPair::from_seed_unchecked(&[7_u8; 32]).unwrap(); + let public = hex_lower(key.public_key().as_ref()); + let key_value = json!({ + "keytype": "ed25519", + "scheme": "ed25519", + "keyval": {"public": public} + }); + let key_id = sha256(&canonical(&key_value)); + let role = json!({"keyids": [key_id.clone()], "threshold": 1}); + let mut keys = Map::new(); + keys.insert(key_id.clone(), key_value); + let root_signed = json!({ + "_type": "root", + "spec_version": "1.0.0", + "consistent_snapshot": false, + "version": 1, + "expires": FUTURE, + "keys": keys, + "roles": { + "root": role.clone(), + "snapshot": role.clone(), + "targets": role.clone(), + "timestamp": role + } + }); + let root = signed_document(&key, &key_id, root_signed); + let root_sha256 = sha256(&root); + + let target = host_target(); + let archive_name = format!("a3s-use-acme-slack-{package_version}-{target}.tar.gz"); + let target_name = + format!("extensions/acme/slack/{package_version}/stable/{target}/{archive_name}"); + let target_sha256 = sha256(&archive); + let mut targets_map = Map::new(); + targets_map.insert( + target_name.clone(), + json!({ + "length": archive.len(), + "hashes": {"sha256": target_sha256}, + "custom": { + "a3s": { + "schemaVersion": 1, + "packageId": "acme/slack", + "version": package_version, + "channel": "stable", + "target": target + } + } + }), + ); + let targets_signed = json!({ + "_type": "targets", + "spec_version": "1.0.0", + "version": metadata_version, + "expires": expires, + "targets": targets_map + }); + let targets = signed_document(&key, &key_id, targets_signed); + let snapshot_signed = json!({ + "_type": "snapshot", + "spec_version": "1.0.0", + "version": metadata_version, + "expires": expires, + "meta": { + "targets.json": { + "version": metadata_version, + "length": targets.len(), + "hashes": {"sha256": sha256(&targets)} + } + } + }); + let snapshot = signed_document(&key, &key_id, snapshot_signed); + let timestamp_signed = json!({ + "_type": "timestamp", + "spec_version": "1.0.0", + "version": metadata_version, + "expires": expires, + "meta": { + "snapshot.json": { + "version": metadata_version, + "length": snapshot.len(), + "hashes": {"sha256": sha256(&snapshot)} + } + } + }); + let timestamp = signed_document(&key, &key_id, timestamp_signed); + + let routes = HashMap::from([ + ("/metadata/root.json".to_string(), root), + ("/metadata/timestamp.json".to_string(), timestamp), + ("/metadata/snapshot.json".to_string(), snapshot), + ("/metadata/targets.json".to_string(), targets), + (format!("/targets/{target_name}"), archive), + ]); + Self { + routes, + root_sha256, + target_name, + target_sha256, + } + } +} + +fn signed_document(key: &Ed25519KeyPair, key_id: &str, signed: Value) -> Vec { + let signature = key.sign(&canonical(&signed)); + serde_json::to_vec(&json!({ + "signatures": [{"keyid": key_id, "sig": hex_lower(signature.as_ref())}], + "signed": signed + })) + .unwrap() +} + +fn canonical(value: &Value) -> Vec { + let mut bytes = Vec::new(); + let mut serializer = + serde_json::Serializer::with_formatter(&mut bytes, CanonicalFormatter::new()); + value.serialize(&mut serializer).unwrap(); + bytes +} + +fn sha256(bytes: &[u8]) -> String { + format!("{:x}", Sha256::digest(bytes)) +} + +fn hex_lower(bytes: &[u8]) -> String { + let mut output = String::with_capacity(bytes.len() * 2); + for byte in bytes { + use std::fmt::Write as _; + let _ = write!(output, "{byte:02x}"); + } + output +} + +pub(crate) fn extension_archive(version: &str) -> Vec { + let manifest = format!( + "extension \"acme/slack\" {{\n schema_version = 1\n version = \"{version}\"\n route = \"slack\"\n actions = [\"read\"]\n\n cli {{\n executable = \"bin/a3s-use-acme-slack\"\n json_output = true\n }}\n}}\n" + ); + let mut bytes = Vec::new(); + { + let encoder = flate2::write::GzEncoder::new(&mut bytes, flate2::Compression::default()); + let mut archive = tar::Builder::new(encoder); + append_tar_file( + &mut archive, + "package/a3s-use-extension.acl", + 0o644, + manifest.as_bytes(), + ); + append_tar_file( + &mut archive, + "package/bin/a3s-use-acme-slack", + 0o755, + b"#!/bin/sh\nprintf 'slack fixture\\n'\n", + ); + archive.finish().unwrap(); + } + bytes +} + +fn append_tar_file(archive: &mut tar::Builder, path: &str, mode: u32, body: &[u8]) { + let mut header = tar::Header::new_gnu(); + header.set_path(path).unwrap(); + header.set_size(body.len() as u64); + header.set_mode(mode); + header.set_cksum(); + archive.append(&header, body).unwrap(); +} + +pub(crate) fn find_subslice(haystack: &[u8], needle: &[u8]) -> Option { + haystack + .windows(needle.len()) + .position(|window| window == needle) +} + +fn host_target() -> &'static str { + match (std::env::consts::OS, std::env::consts::ARCH) { + ("macos", "aarch64") => "darwin-arm64", + ("macos", "x86_64") => "darwin-x86_64", + ("linux", "aarch64") => "linux-arm64", + ("linux", "x86_64") => "linux-x86_64", + ("windows", "x86_64") => "windows-x86_64", + (os, arch) => panic!("unsupported TUF test target {os}-{arch}"), + } +} + +pub(crate) struct TestServer { + base_url: String, + routes: Arc>>>, + requests: Arc>>, + stop: Arc, + thread: Option>, +} + +impl TestServer { + pub(crate) fn start(routes: HashMap>) -> Self { + let listener = TcpListener::bind("127.0.0.1:0").unwrap(); + listener.set_nonblocking(true).unwrap(); + let base_url = format!("http://{}/", listener.local_addr().unwrap()); + let routes = Arc::new(Mutex::new(routes)); + let requests = Arc::new(Mutex::new(Vec::new())); + let stop = Arc::new(AtomicBool::new(false)); + let thread_routes = Arc::clone(&routes); + let thread_requests = Arc::clone(&requests); + let thread_stop = Arc::clone(&stop); + let thread = std::thread::spawn(move || { + while !thread_stop.load(Ordering::Relaxed) { + match listener.accept() { + Ok((stream, _)) => { + stream.set_nonblocking(false).unwrap(); + let routes = Arc::clone(&thread_routes); + let requests = Arc::clone(&thread_requests); + std::thread::spawn(move || serve(stream, &routes, &requests)); + } + Err(error) if error.kind() == std::io::ErrorKind::WouldBlock => { + std::thread::sleep(Duration::from_millis(5)); + } + Err(_) => break, + } + } + }); + Self { + base_url, + routes, + requests, + stop, + thread: Some(thread), + } + } + + pub(crate) fn base_url(&self) -> &str { + &self.base_url + } + + pub(crate) fn requests(&self) -> Vec { + self.requests.lock().unwrap().clone() + } + + pub(crate) fn clear_requests(&self) { + self.requests.lock().unwrap().clear(); + } + + pub(crate) fn replace_routes(&self, routes: HashMap>) { + *self.routes.lock().unwrap() = routes; + } +} + +impl Drop for TestServer { + fn drop(&mut self) { + self.stop.store(true, Ordering::Relaxed); + if let Some(thread) = self.thread.take() { + let _ = thread.join(); + } + } +} + +fn serve( + mut stream: TcpStream, + routes: &Mutex>>, + requests: &Mutex>, +) { + let _ = stream.set_read_timeout(Some(Duration::from_secs(2))); + let mut buffer = [0_u8; 8192]; + let Ok(size) = stream.read(&mut buffer) else { + return; + }; + let request = String::from_utf8_lossy(&buffer[..size]); + let path = request + .lines() + .next() + .and_then(|line| line.split_whitespace().nth(1)) + .unwrap_or("/") + .to_string(); + requests.lock().unwrap().push(path.clone()); + let body = routes.lock().unwrap().get(&path).cloned(); + let (status, body) = body + .as_deref() + .map(|body| ("200 OK", body)) + .unwrap_or(("404 Not Found", b"not found")); + let header = format!( + "HTTP/1.1 {status}\r\nContent-Length: {}\r\nConnection: close\r\n\r\n", + body.len() + ); + if stream.write_all(header.as_bytes()).is_ok() && stream.write_all(body).is_ok() { + let _ = stream.flush(); + let _ = stream.shutdown(Shutdown::Write); + } +} diff --git a/docs/architecture.md b/docs/architecture.md index c165d7b5..1408dc9d 100644 --- a/docs/architecture.md +++ b/docs/architecture.md @@ -70,6 +70,14 @@ Consumers read `extension snapshot` for the current projection or long-poll `extension watch --after-generation ` for a later generation. No daemon, custom RPC protocol, `dlopen`, or restart is required. +Explicit local sources may be directories or bounded `.tar.gz`, `.tgz`, and +`.zip` archives. Archive extraction runs off the async executor, accepts one +manifest-rooted package, preserves executable permissions, and rejects links, +path traversal, duplicate entries, unsupported file types, and expansion beyond +the package limits before lifecycle activation begins. Standard bounded macOS +AppleDouble sidecars are ignored rather than installed; they still count toward +the archive entry and expanded-byte limits. + ### Unified capability projection Resident Code hosts do not need separate discovery paths for built-in and @@ -500,6 +508,11 @@ Implemented: vision adapters, standard MCP annotations/output schemas, and a release-packaged content-bound Skill that projects to `mcp__use_ocr__*` in A3S Code. +17. TUF-verified remote extension registries with pinned bootstrap roots, + expiration and rollback enforcement, exact review/apply plans, and signed + provenance receipts. Registry upgrades restore the recorded source and + channel, reject identity drift and version downgrades, and converge before + payload download when the installed signed target is already current. Next: @@ -513,5 +526,6 @@ Next: with the same runtime guarantees as macOS and Linux. Windows compilation, CLI/MCP schemas, packaged assets, and non-runtime tests remain continuously checked in CI meanwhile. -4. Signed remote extension publishers. External publisher infrastructure is - independent of the built-in Browser compatibility contract. +4. Production publication for the official A3S extension registry, including + an offline-held root-key policy and release automation. The client does not + substitute a placeholder or generated key for that operational trust root. diff --git a/src/cli.rs b/src/cli.rs index 23bad00f..42d387f0 100644 --- a/src/cli.rs +++ b/src/cli.rs @@ -6,7 +6,8 @@ use crate::capability_registry::{ use crate::extension_cli::{ extension_capabilities, extension_disable, extension_enable, extension_inspect, extension_list, extension_snapshot, extension_watch, external_component_value, external_package_id, - install_extension, installed_extension, installed_extensions, uninstall_extension, + install_extension, install_remote_extension, installed_extension, installed_extensions, + uninstall_extension, }; use std::time::Duration; @@ -451,15 +452,78 @@ async fn component_install(args: &[String]) -> UseResult { format!("Unknown delegated component '{id}'."), )); }; - let source = option_argument(args, "--from")? - .ok_or_else(|| usage_error("external extension install requires --from "))?; - let result = install_extension( - package_id, - std::path::Path::new(source), - args.iter().any(|argument| argument == "--force"), - args.iter().any(|argument| argument == "--allow-unsigned"), - ) - .await?; + let source = option_argument(args, "--from")?; + let registry_name = option_argument(args, "--registry-name")?; + let registry_url = option_argument(args, "--registry-url")?; + let trust_root = option_argument(args, "--trust-root")?; + let trusted_root = option_argument(args, "--trusted-root")?; + let version = option_argument(args, "--version")?; + let channel = option_argument(args, "--channel")?.unwrap_or("stable"); + let expected_plan = option_argument(args, "--registry-plan-digest")?; + let force = args.iter().any(|argument| argument == "--force"); + let allow_unsigned = args.iter().any(|argument| argument == "--allow-unsigned"); + let remote_requested = registry_name.is_some() + || registry_url.is_some() + || trust_root.is_some() + || trusted_root.is_some() + || version.is_some() + || expected_plan.is_some() + || option_argument(args, "--channel")?.is_some(); + let result = if let Some(source) = source { + if remote_requested { + return Err(usage_error( + "--from cannot be combined with signed registry options", + )); + } + install_extension( + package_id, + std::path::Path::new(source), + force, + allow_unsigned, + ) + .await? + } else { + if allow_unsigned { + return Err(usage_error( + "--allow-unsigned is valid only with an explicit local --from package", + )); + } + let registry_name = registry_name + .ok_or_else(|| usage_error("remote extension install requires --registry-name"))?; + let registry_url = registry_url + .ok_or_else(|| usage_error("remote extension install requires --registry-url"))?; + let trust_root = trust_root + .ok_or_else(|| usage_error("remote extension install requires --trust-root"))?; + let trusted_root = trusted_root + .map(|path| { + let path = std::path::PathBuf::from(path); + if path.is_absolute() { + Ok(path) + } else { + std::env::current_dir() + .map(|directory| directory.join(path)) + .map_err(|error| { + UseError::new( + "use.extension.registry_path_invalid", + format!("Failed to resolve the trusted root path: {error}"), + ) + }) + } + }) + .transpose()?; + install_remote_extension( + package_id, + registry_name, + registry_url, + trust_root, + trusted_root.as_deref(), + version, + channel, + expected_plan, + force, + ) + .await? + }; Ok(CommandOutput::success( if result.changed { format!("Installed extension '{}'.", result.extension.package_id) @@ -958,9 +1022,16 @@ fn validate_component_install_options(args: &[String]) -> UseResult<()> { while index < args.len() { match args[index].as_str() { "--json" | "--force" | "--allow-unsigned" => index += 1, - "--from" => { + "--from" + | "--registry-name" + | "--registry-url" + | "--trust-root" + | "--trusted-root" + | "--version" + | "--channel" + | "--registry-plan-digest" => { if args.get(index + 1).is_none() { - return Err(usage_error("--from requires a value")); + return Err(usage_error(format!("{} requires a value", args[index]))); } index += 2; } diff --git a/src/extension_cli.rs b/src/extension_cli.rs index 4e8afb64..ee5b30e2 100644 --- a/src/extension_cli.rs +++ b/src/extension_cli.rs @@ -14,6 +14,8 @@ pub(crate) struct ExtensionView { pub enabled: bool, pub package_root: PathBuf, pub surfaces: Vec<&'static str>, + pub trust: &'static str, + pub registry: Option, pub manifest: serde_json::Value, } @@ -196,7 +198,8 @@ pub(crate) fn external_component_value( "route": extension.route, "enabled": extension.enabled, "surfaces": extension.surfaces, - "trust": "local-explicit" + "trust": extension.trust, + "registry": extension.registry }) } @@ -209,7 +212,8 @@ fn extension_value(extension: &ExtensionView) -> serde_json::Value { "enabled": extension.enabled, "packageRoot": extension.package_root, "surfaces": extension.surfaces, - "trust": "local-explicit" + "trust": extension.trust, + "registry": extension.registry }) } @@ -254,6 +258,42 @@ pub(crate) async fn install_extension( }) } +#[cfg(feature = "extensions")] +#[allow(clippy::too_many_arguments)] +pub(crate) async fn install_remote_extension( + package_id: &str, + registry_name: &str, + registry_url: &str, + trust_root: &str, + trusted_root_path: Option<&Path>, + version: Option<&str>, + channel: &str, + expected_plan_digest: Option<&str>, + force: bool, +) -> UseResult { + let paths = a3s_use_extension::ExtensionPaths::from_env()?; + let registry = a3s_use_extension::TrustedRegistry::new( + registry_name, + registry_url, + trust_root, + trusted_root_path.map(Path::to_path_buf), + paths.tuf_datastore(registry_name), + )?; + let result = crate::extension_host::install_remote( + package_id, + ®istry, + version, + channel, + expected_plan_digest, + force, + ) + .await?; + Ok(ExtensionInstallView { + changed: result.changed, + extension: extension_view(result.extension)?, + }) +} + #[cfg(not(feature = "extensions"))] pub(crate) async fn install_extension( _package_id: &str, @@ -264,6 +304,22 @@ pub(crate) async fn install_extension( Err(extensions_disabled()) } +#[cfg(not(feature = "extensions"))] +#[allow(clippy::too_many_arguments)] +pub(crate) async fn install_remote_extension( + _package_id: &str, + _registry_name: &str, + _registry_url: &str, + _trust_root: &str, + _trusted_root_path: Option<&Path>, + _version: Option<&str>, + _channel: &str, + _expected_plan_digest: Option<&str>, + _force: bool, +) -> UseResult { + Err(extensions_disabled()) +} + #[cfg(feature = "extensions")] pub(crate) async fn uninstall_extension(package_id: &str) -> UseResult { let result = crate::extension_host::uninstall(package_id).await?; @@ -368,6 +424,22 @@ async fn watch_registry( #[cfg(feature = "extensions")] fn extension_view(extension: a3s_use_extension::InstalledExtension) -> UseResult { let surfaces = extension.surfaces(); + let trust = match extension.receipt.trust { + a3s_use_extension::ExtensionTrust::LocalExplicit => "local-explicit", + a3s_use_extension::ExtensionTrust::RegistryTuf => "registry-tuf", + }; + let registry = extension + .receipt + .registry + .as_ref() + .map(serde_json::to_value) + .transpose() + .map_err(|error| { + UseError::new( + "use.extension.receipt_invalid", + format!("Failed to encode the extension registry provenance: {error}"), + ) + })?; let manifest = serde_json::to_value(&extension.manifest).map_err(|error| { UseError::new( "use.extension.manifest_invalid", @@ -382,6 +454,8 @@ fn extension_view(extension: a3s_use_extension::InstalledExtension) -> UseResult enabled: extension.receipt.enabled, package_root: extension.receipt.package_root, surfaces, + trust, + registry, manifest, }) } diff --git a/src/extension_host.rs b/src/extension_host.rs index b2e845e7..a2fc60bf 100644 --- a/src/extension_host.rs +++ b/src/extension_host.rs @@ -3,7 +3,7 @@ use std::path::Path; use a3s_use_core::{UseError, UseResult}; use a3s_use_extension::{ ActivationResult, ExtensionRegistry, ExtensionRegistrySnapshot, InstallOptions, InstallResult, - InstalledExtension, UninstallResult, + InstalledExtension, TrustedRegistry, UninstallResult, }; use std::time::Duration; @@ -33,6 +33,26 @@ pub async fn install( .await } +pub async fn install_remote( + package_id: &str, + registry: &TrustedRegistry, + version: Option<&str>, + channel: &str, + expected_plan_digest: Option<&str>, + force: bool, +) -> UseResult { + ExtensionRegistry::from_env()? + .install_remote( + package_id, + registry, + version, + channel, + expected_plan_digest, + force, + ) + .await +} + pub async fn uninstall(package_id: &str) -> UseResult { ExtensionRegistry::from_env()?.uninstall(package_id).await } diff --git a/tests/extension_archives.rs b/tests/extension_archives.rs new file mode 100644 index 00000000..c3d81e32 --- /dev/null +++ b/tests/extension_archives.rs @@ -0,0 +1,102 @@ +#![cfg(all(unix, feature = "extensions"))] + +use std::fs::File; +use std::os::unix::fs::PermissionsExt; +use std::process::Command; + +fn binary() -> &'static str { + env!("CARGO_BIN_EXE_a3s-use") +} + +#[test] +fn archived_extension_installs_dispatches_and_uninstalls_through_the_cli() { + let temp = tempfile::tempdir().unwrap(); + let package = temp.path().join("package"); + std::fs::create_dir_all(package.join("bin")).unwrap(); + std::fs::write( + package.join("a3s-use-extension.acl"), + r#"extension "acme/slack" { + schema_version = 1 + version = "1.0.0" + route = "slack" + actions = ["read"] + + cli { + executable = "bin/a3s-use-acme-slack" + json_output = true + } +} +"#, + ) + .unwrap(); + let executable = package.join("bin/a3s-use-acme-slack"); + std::fs::write( + &executable, + "#!/bin/sh\nprintf '%s\\n' \"$A3S_USE_EXTENSION_ID\"\nprintf '%s\\n' \"$*\"\nexit 7\n", + ) + .unwrap(); + std::fs::set_permissions(&executable, std::fs::Permissions::from_mode(0o755)).unwrap(); + + let archive = temp.path().join("acme-slack.tar.gz"); + let file = File::create(&archive).unwrap(); + let encoder = flate2::write::GzEncoder::new(file, flate2::Compression::default()); + let mut builder = tar::Builder::new(encoder); + builder.append_dir_all("package", &package).unwrap(); + builder.finish().unwrap(); + drop(builder); + + let home = temp.path().join("home"); + let installed = Command::new(binary()) + .args([ + "component", + "install", + "acme/slack", + "--from", + archive.to_str().unwrap(), + "--allow-unsigned", + "--json", + ]) + .env("A3S_USE_HOME", &home) + .output() + .unwrap(); + assert!( + installed.status.success(), + "status: {}\nstdout: {}\nstderr: {}", + installed.status, + String::from_utf8_lossy(&installed.stdout), + String::from_utf8_lossy(&installed.stderr) + ); + let installed_json: serde_json::Value = serde_json::from_slice(&installed.stdout).unwrap(); + assert_eq!(installed_json["data"]["component"]["id"], "acme/slack"); + + let delegated = Command::new(binary()) + .args(["slack", "channels", "list", "--json"]) + .env("A3S_USE_HOME", &home) + .output() + .unwrap(); + assert_eq!(delegated.status.code(), Some(7)); + assert_eq!( + String::from_utf8(delegated.stdout).unwrap(), + "acme/slack\nchannels list --json\n" + ); + + let removed = Command::new(binary()) + .args(["component", "uninstall", "acme/slack", "--json"]) + .env("A3S_USE_HOME", &home) + .output() + .unwrap(); + assert!(removed.status.success(), "{removed:?}"); + let removed_json: serde_json::Value = serde_json::from_slice(&removed.stdout).unwrap(); + assert_eq!(removed_json["data"]["changed"], true); + + let listed = Command::new(binary()) + .args(["extension", "list", "--json"]) + .env("A3S_USE_HOME", &home) + .output() + .unwrap(); + assert!(listed.status.success(), "{listed:?}"); + let listed_json: serde_json::Value = serde_json::from_slice(&listed.stdout).unwrap(); + assert_eq!(listed_json["data"]["extensions"], serde_json::json!([])); + assert!(!home.join("data/extensions/acme/slack").exists()); + assert!(!home.join("state/extensions/acme/slack.json").exists()); +} diff --git a/tests/remote_extension_cli.rs b/tests/remote_extension_cli.rs new file mode 100644 index 00000000..74167a76 --- /dev/null +++ b/tests/remote_extension_cli.rs @@ -0,0 +1,183 @@ +#![cfg(feature = "extensions")] + +use std::process::{Command, Output}; + +use a3s_use_extension::{prepare_remote_package, ResolvedRemotePackage, TrustedRegistry}; + +#[path = "../crates/extension/src/tuf_test_support.rs"] +mod tuf_test_support; + +use tuf_test_support::{extension_archive, TestRepository, TestServer, FUTURE, PACKAGE_VERSION}; + +fn binary() -> &'static str { + env!("CARGO_BIN_EXE_a3s-use") +} + +#[tokio::test] +async fn signed_registry_install_uses_reviewed_target_and_reports_tuf_provenance() { + let repository = TestRepository::new(extension_archive(PACKAGE_VERSION), 1, FUTURE); + let server = TestServer::start(repository.routes.clone()); + let temp = tempfile::tempdir().unwrap(); + let trusted = TrustedRegistry::new( + "fixture", + server.base_url(), + &repository.root_sha256, + None, + temp.path().join("review-state"), + ) + .unwrap(); + let reviewed = prepare_remote_package(&trusted, "acme/slack", None, "stable", None) + .await + .unwrap(); + let plan_digest = reviewed.resolved().plan_digest().unwrap(); + drop(reviewed); + assert_no_target_request(&server); + + let home = temp.path().join("home"); + let installed = registry_install(&server, &repository, &home, Some(&plan_digest), &[]); + assert!(installed.status.success(), "{installed:?}"); + let installed_json = json(&installed); + assert_eq!(installed_json["data"]["changed"], true); + assert_eq!(installed_json["data"]["component"]["trust"], "registry-tuf"); + assert_eq!( + installed_json["data"]["component"]["registry"]["registryName"], + "fixture" + ); + assert_eq!( + installed_json["data"]["component"]["registry"]["sha256"], + repository.target_sha256 + ); + assert_eq!( + server + .requests() + .iter() + .filter(|request| request.starts_with("/targets/")) + .count(), + 1 + ); + + let receipt: serde_json::Value = serde_json::from_slice( + &std::fs::read(home.join("state/extensions/acme/slack.json")).unwrap(), + ) + .unwrap(); + assert_eq!(receipt["trust"], "registry-tuf"); + let provenance: ResolvedRemotePackage = + serde_json::from_value(receipt["registry"].clone()).unwrap(); + assert_eq!(provenance.plan_digest().unwrap(), plan_digest); + + let inspected = Command::new(binary()) + .args(["extension", "inspect", "acme/slack", "--json"]) + .env("A3S_USE_HOME", &home) + .output() + .unwrap(); + assert!(inspected.status.success(), "{inspected:?}"); + let inspected = json(&inspected); + assert_eq!(inspected["data"]["extension"]["trust"], "registry-tuf"); + assert_eq!( + inspected["data"]["extension"]["registry"]["targetName"], + repository.target_name + ); + + let second = registry_install(&server, &repository, &home, Some(&plan_digest), &[]); + assert!(second.status.success(), "{second:?}"); + assert_eq!(json(&second)["data"]["changed"], false); +} + +#[test] +fn registry_plan_mismatch_fails_before_target_download() { + let repository = TestRepository::new(extension_archive(PACKAGE_VERSION), 1, FUTURE); + let server = TestServer::start(repository.routes.clone()); + let temp = tempfile::tempdir().unwrap(); + let output = registry_install( + &server, + &repository, + &temp.path().join("home"), + Some(&"0".repeat(64)), + &[], + ); + + assert!(!output.status.success(), "{output:?}"); + assert_eq!( + json(&output)["error"]["code"], + "use.extension.registry_plan_mismatch" + ); + assert_no_target_request(&server); +} + +#[test] +fn registry_install_rejects_unsigned_and_local_source_combinations() { + let repository = TestRepository::new(extension_archive(PACKAGE_VERSION), 1, FUTURE); + let server = TestServer::start(repository.routes.clone()); + let temp = tempfile::tempdir().unwrap(); + let home = temp.path().join("home"); + + let unsigned = registry_install(&server, &repository, &home, None, &["--allow-unsigned"]); + assert!(!unsigned.status.success(), "{unsigned:?}"); + assert_eq!(json(&unsigned)["error"]["code"], "use.cli.invalid_usage"); + + let local = Command::new(binary()) + .args([ + "component", + "install", + "acme/slack", + "--from", + temp.path().to_str().unwrap(), + "--allow-unsigned", + "--registry-name", + "fixture", + "--json", + ]) + .env("A3S_USE_HOME", &home) + .output() + .unwrap(); + assert!(!local.status.success(), "{local:?}"); + assert_eq!(json(&local)["error"]["code"], "use.cli.invalid_usage"); + assert!(server.requests().is_empty()); +} + +fn registry_install( + server: &TestServer, + repository: &TestRepository, + home: &std::path::Path, + plan_digest: Option<&str>, + extra: &[&str], +) -> Output { + let mut command = Command::new(binary()); + command.args([ + "component", + "install", + "acme/slack", + "--registry-name", + "fixture", + "--registry-url", + server.base_url(), + "--trust-root", + &repository.root_sha256, + ]); + if let Some(plan_digest) = plan_digest { + command.args(["--registry-plan-digest", plan_digest]); + } + command + .args(extra) + .arg("--json") + .env("A3S_USE_HOME", home) + .output() + .unwrap() +} + +fn json(output: &Output) -> serde_json::Value { + serde_json::from_slice(&output.stdout).unwrap_or_else(|error| { + panic!( + "invalid JSON output ({error}): stdout={:?}, stderr={:?}", + String::from_utf8_lossy(&output.stdout), + String::from_utf8_lossy(&output.stderr) + ) + }) +} + +fn assert_no_target_request(server: &TestServer) { + assert!(server + .requests() + .iter() + .all(|request| !request.starts_with("/targets/"))); +} From fd8649bd20f90613863ba74ae6734066336c9d8a Mon Sep 17 00:00:00 2001 From: RoyLin Date: Sun, 19 Jul 2026 10:23:41 +0800 Subject: [PATCH 3/9] feat(science): add typed life-science extension --- Cargo.lock | 18 + Cargo.toml | 1 + README.md | 34 ++ crates/science/Cargo.toml | 33 ++ crates/science/DATA_SOURCES.md | 31 ++ crates/science/README.md | 85 ++++ crates/science/UPSTREAM.md | 26 ++ crates/science/package/a3s-use-extension.acl | 21 + .../package/skills/a3s-use-science/SKILL.md | 47 ++ crates/science/scripts/package.sh | 25 + crates/science/src/biorxiv.rs | 293 ++++++++++++ crates/science/src/chembl.rs | 270 +++++++++++ crates/science/src/cli.rs | 370 +++++++++++++++ crates/science/src/client.rs | 365 +++++++++++++++ crates/science/src/clinical_trials.rs | 312 +++++++++++++ crates/science/src/ensembl.rs | 203 ++++++++ crates/science/src/lib.rs | 24 + crates/science/src/main.rs | 39 ++ crates/science/src/mcp.rs | 434 ++++++++++++++++++ crates/science/src/models.rs | 130 ++++++ crates/science/src/pubmed.rs | 248 ++++++++++ crates/science/tests/integration.rs | 318 +++++++++++++ docs/architecture.md | 8 + 23 files changed, 3335 insertions(+) create mode 100644 crates/science/Cargo.toml create mode 100644 crates/science/DATA_SOURCES.md create mode 100644 crates/science/README.md create mode 100644 crates/science/UPSTREAM.md create mode 100644 crates/science/package/a3s-use-extension.acl create mode 100644 crates/science/package/skills/a3s-use-science/SKILL.md create mode 100755 crates/science/scripts/package.sh create mode 100644 crates/science/src/biorxiv.rs create mode 100644 crates/science/src/chembl.rs create mode 100644 crates/science/src/cli.rs create mode 100644 crates/science/src/client.rs create mode 100644 crates/science/src/clinical_trials.rs create mode 100644 crates/science/src/ensembl.rs create mode 100644 crates/science/src/lib.rs create mode 100644 crates/science/src/main.rs create mode 100644 crates/science/src/mcp.rs create mode 100644 crates/science/src/models.rs create mode 100644 crates/science/src/pubmed.rs create mode 100644 crates/science/tests/integration.rs diff --git a/Cargo.lock b/Cargo.lock index 14cb1eb3..d4f96f2f 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -135,6 +135,24 @@ dependencies = [ "zip", ] +[[package]] +name = "a3s-use-science" +version = "0.1.1" +dependencies = [ + "a3s-use-core", + "a3s-use-extension", + "axum", + "clap", + "reqwest", + "rmcp", + "schemars", + "serde", + "serde_json", + "tempfile", + "tokio", + "url", +] + [[package]] name = "adler2" version = "2.0.1" diff --git a/Cargo.toml b/Cargo.toml index a136a13b..9b807729 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -5,6 +5,7 @@ members = [ "crates/browser-driver", "crates/office", "crates/extension", + "crates/science", ] resolver = "2" diff --git a/README.md b/README.md index 9b8376a1..90ff91a5 100644 --- a/README.md +++ b/README.md @@ -129,6 +129,8 @@ Every domain argument accepted by `a3s use ...` can also be passed directly to safe Word, Spreadsheet, Presentation, native MCP, and compatibility workflows - **External Domains**: Install process-isolated packages that expose any useful combination of CLI, MCP, and Skill surfaces +- **Reference Science Toolkit**: Query PubMed, ChEMBL, ClinicalTrials.gov, + bioRxiv, and Ensembl through one typed read-only extension - **Hot-Plug Discovery**: Publish immutable generation/revision snapshots so a resident host can add, replace, or remove live capabilities without restarting - **Content-Bound Skills**: Project an absolute package path and lowercase @@ -149,6 +151,7 @@ Every domain argument accepted by `a3s use ...` can also be passed directly to | Browser | Built in | Full Browser vocabulary | A3S Use standard MCP server | Six packaged Browser Skills | A3S Use | | Office | Built in | Stable Office vocabulary | Typed native preview plus OfficeCLI compatibility server | Packaged `a3s-use-office` Skill | A3S Use native engine; OfficeCLI compatibility in 0.1.x | | Box | Reserved built-in route | Native A3S Box vocabulary | — | — | Umbrella A3S CLI | +| Science | External `a3s/science` package | Source-specific retrieval commands | 13 typed `science_*` tools | One research workflow Skill | Science extension process | | External domain | Installed extension | Optional native executable | Optional standard MCP server | Optional `SKILL.md` | Extension package plus A3S Use lifecycle | The Box route is component-backed. The umbrella CLI resolves its authoritative @@ -179,6 +182,7 @@ A compiled command surface is not proof that its provider is installed. Use | `a3s-use-browser-driver` | Complete interactive Browser CLI, MCP tools, Skills, Dashboard, and compatibility runtime | | `a3s-use-office` | Native OOXML foundation, typed Office operations, and compatibility lifecycle | | `a3s-use-extension` | A3S ACL manifest model, package registry, leases, and native surface descriptors | +| `a3s-use-science` | Typed public life-science APIs, CLI, MCP tools, and extension package assets | | `a3s-use` | Facade library, standalone CLI host, capability projection, and MCP entry points | ## Quick Start @@ -1586,6 +1590,36 @@ compatibility response can return See [Native Office Engine](docs/native-office.md) for the complete requirements, compatibility scope, safety invariants, delivery gates, and migration plan. +## Science Toolkit + +The repository includes `a3s-use-science` as a reference external extension, +not as another built-in route. Its process exposes one typed Rust client as 13 +read-only MCP tools plus source-specific CLI commands for PubMed, ChEMBL, +ClinicalTrials.gov, bioRxiv, and Ensembl. + +Build a local package into a new directory and install it explicitly: + +```bash +./crates/science/scripts/package.sh /tmp/a3s-use-science-package +a3s install use/a3s/science \ + --from /tmp/a3s-use-science-package \ + --allow-unsigned + +export A3S_SCIENCE_CONTACT_EMAIL=researcher@example.org +a3s use science pubmed search "single-cell atlas" --limit 10 --json +a3s use science ensembl lookup homo_sapiens TP53 --json +a3s use mcp serve a3s/science +``` + +The same package can be archived as `.tar.gz`, `.tgz`, or `.zip` and installed +through the explicit local-package flow. Local packages require +`--allow-unsigned`; use them only after review. PubMed requires the contact +email, while `NCBI_API_KEY` is optional. See the +[Science crate](crates/science/README.md), its +[data-source notice](crates/science/DATA_SOURCES.md), and +[clean-room provenance](crates/science/UPSTREAM.md) for the full command set, +data egress, limits, and interpretation boundaries. + ## External Extensions External Use domains stay behind process boundaries. A package contains an diff --git a/crates/science/Cargo.toml b/crates/science/Cargo.toml new file mode 100644 index 00000000..6e99d89d --- /dev/null +++ b/crates/science/Cargo.toml @@ -0,0 +1,33 @@ +[package] +name = "a3s-use-science" +version.workspace = true +edition.workspace = true +license.workspace = true +repository.workspace = true +authors.workspace = true +rust-version.workspace = true +description = "Typed life-science data retrieval for A3S Use" + +[lib] +name = "a3s_use_science" +path = "src/lib.rs" + +[[bin]] +name = "a3s-use-science" +path = "src/main.rs" + +[dependencies] +a3s-use-core = { version = "0.1.1", path = "../core" } +clap.workspace = true +reqwest = { workspace = true, features = ["json"] } +rmcp.workspace = true +schemars.workspace = true +serde.workspace = true +serde_json.workspace = true +tokio.workspace = true +url.workspace = true + +[dev-dependencies] +a3s-use-extension = { version = "0.1.1", path = "../extension" } +axum.workspace = true +tempfile.workspace = true diff --git a/crates/science/DATA_SOURCES.md b/crates/science/DATA_SOURCES.md new file mode 100644 index 00000000..4c37b0e2 --- /dev/null +++ b/crates/science/DATA_SOURCES.md @@ -0,0 +1,31 @@ +# Science Data Sources + +The Science extension sends user-supplied search terms and identifiers over +HTTPS to public third-party services. A query can reveal research interests; +do not submit confidential, patient-identifying, controlled, or unpublished +information unless the applicable policy and upstream terms permit it. + +| Source | Endpoint | Data returned | Local credential | +| --- | --- | --- | --- | +| PubMed / NCBI E-utilities | `eutils.ncbi.nlm.nih.gov` | Citation summaries and identifiers | Contact email required; API key optional | +| ChEMBL | `www.ebi.ac.uk/chembl` | Molecules, targets, and activities | None | +| ClinicalTrials.gov | `clinicaltrials.gov/api/v2` | Public study protocol records | None | +| bioRxiv | `api.biorxiv.org` | Public preprint metadata | None | +| Ensembl REST | `rest.ensembl.org` | Public gene and homology records | None | + +The extension applies per-source request pacing, bounded result limits, a +30-second default request timeout, and bounded upstream error bodies. bioRxiv +free-text filtering scans at most 500 records per command. PubMed requests +identify `a3s-use-science` and include the configured contact email in line +with NCBI guidance. + +Upstream services remain authoritative for licenses, terms, retention, +availability, update cadence, and record interpretation. Their schemas and +content can change independently of A3S. A successful response means only that +the public API returned a record; it does not establish scientific validity, +peer review, clinical suitability, or regulatory approval. + +Always preserve source identifiers and retrieval context. Label bioRxiv +records as preprints, verify consequential conclusions against the underlying +publication or protocol, and do not use this toolkit as a substitute for +medical, safety, ethics, or regulatory review. diff --git a/crates/science/README.md b/crates/science/README.md new file mode 100644 index 00000000..2259e755 --- /dev/null +++ b/crates/science/README.md @@ -0,0 +1,85 @@ +# A3S Use Science + +`a3s-use-science` is a process-isolated, read-only life-science extension for +A3S Use. It provides one typed asynchronous Rust client and projects the same +operations through a native CLI and a standard MCP server. + +The initial toolkit covers: + +| Source | Operations | +| --- | --- | +| PubMed | Search article summaries; retrieve a PMID | +| ChEMBL | Search molecules and targets; retrieve molecules and activities | +| ClinicalTrials.gov | Search studies; retrieve an NCT record | +| bioRxiv | Search a bounded date range; retrieve a DOI | +| Ensembl | Look up a gene; retrieve orthologs | + +All operations are retrieval-only. The crate does not copy implementation code +from upstream skill collections and does not run their Python environments. +See [UPSTREAM.md](UPSTREAM.md) for the inspiration, reviewed revision, and +clean-room boundary. + +## Configuration + +Set a contact email before using PubMed, as requested by NCBI E-utilities: + +```bash +export A3S_SCIENCE_CONTACT_EMAIL=researcher@example.org +export NCBI_API_KEY=optional-ncbi-key +``` + +`NCBI_API_KEY` is optional. The other sources currently use public endpoints +without credentials. See [DATA_SOURCES.md](DATA_SOURCES.md) for network, +provenance, and usage considerations. + +## CLI + +Build and run from the A3S Use workspace: + +```bash +cargo build -p a3s-use-science +./target/debug/a3s-use-science doctor --json +./target/debug/a3s-use-science pubmed search "single-cell atlas" --limit 10 --json +./target/debug/a3s-use-science chembl get-molecule CHEMBL25 --json +./target/debug/a3s-use-science clinical-trials search glioblastoma --status RECRUITING --json +./target/debug/a3s-use-science biorxiv search --from 2026-01-01 --to 2026-01-31 --json +./target/debug/a3s-use-science ensembl lookup homo_sapiens BRCA1 --json +``` + +Every `--json` invocation returns one versioned CLI document. Without +`--json`, commands print the retrieved typed value as readable JSON. + +## Standard MCP + +Run the extension's stdio MCP server directly with: + +```bash +./target/debug/a3s-use-science serve --mcp +``` + +After packaging and installing the extension, the A3S host route is: + +```bash +a3s use mcp serve a3s/science +``` + +The server exposes 13 source-specific `science_*` tools. It does not introduce +an A3S-specific RPC envelope or combine unrelated source vocabularies into a +generic execute action. + +## Package + +Create a local extension directory at a new path: + +```bash +./crates/science/scripts/package.sh /tmp/a3s-use-science-package +a3s install use/a3s/science \ + --from /tmp/a3s-use-science-package \ + --allow-unsigned +a3s use science doctor --json +``` + +The script refuses to overwrite an existing output directory. The package may +also be archived as `.tar.gz`, `.tgz`, or `.zip` and passed directly to +`--from`. Local directories and archives require explicit `--allow-unsigned` +trust; a signed remote distribution channel remains roadmap work. diff --git a/crates/science/UPSTREAM.md b/crates/science/UPSTREAM.md new file mode 100644 index 00000000..74f9b447 --- /dev/null +++ b/crates/science/UPSTREAM.md @@ -0,0 +1,26 @@ +# Upstream Inspiration and Clean-Room Boundary + +The public capability inventory in +[`baifan-wang/skills/claude-science`](https://github.com/baifan-wang/skills/tree/main/claude-science) +inspired the source selection and agent workflow for this extension. The +inventory was reviewed at commit +`2b61d890c5ba50570717599b16d34514458b3955` on 2026-07-17. + +That repository describes a much larger collection of data tools, model +workflows, compute integrations, and scientific Skills. Its components carry +component-specific licensing rather than one clearly stated project-level +license. Consequently, `a3s-use-science` is an independent clean-room Rust +implementation: + +- no Python source, JSON schema, prompt, test fixture, model wrapper, or other + implementation artifact is copied or distributed; +- this initial package implements a smaller source-specific retrieval surface + and does not claim command, MCP-tool, or output compatibility; +- Ensembl access uses the documented Ensembl REST API, not the upstream + collection's BioMart implementation; +- the extension's code and package assets are licensed with A3S Use under MIT. + +Public database names and documented HTTP contracts are factual integration +points, not bundled upstream software. Each remote data service retains its own +terms, licenses, and attribution requirements; see +[DATA_SOURCES.md](DATA_SOURCES.md). diff --git a/crates/science/package/a3s-use-extension.acl b/crates/science/package/a3s-use-extension.acl new file mode 100644 index 00000000..97a5fe49 --- /dev/null +++ b/crates/science/package/a3s-use-extension.acl @@ -0,0 +1,21 @@ +extension "a3s/science" { + schema_version = 1 + version = "0.1.1" + route = "science" + actions = ["read"] + + cli { + executable = "bin/a3s-use-science" + json_output = true + } + + mcp { + executable = "bin/a3s-use-science" + args = ["serve", "--mcp"] + transport = "stdio" + } + + skill { + path = "skills/a3s-use-science/SKILL.md" + } +} diff --git a/crates/science/package/skills/a3s-use-science/SKILL.md b/crates/science/package/skills/a3s-use-science/SKILL.md new file mode 100644 index 00000000..6a6449b1 --- /dev/null +++ b/crates/science/package/skills/a3s-use-science/SKILL.md @@ -0,0 +1,47 @@ +--- +name: a3s-use-science +description: Retrieve and cross-check public biomedical evidence from PubMed, ChEMBL, ClinicalTrials.gov, bioRxiv, and Ensembl. Use for literature searches, preprint checks, compound and target research, trial discovery, gene lookup, and ortholog analysis through A3S Use. +allowed-tools: Bash(a3s:*) +--- + +# A3S Use Science + +Use the host surface that is already available: + +- In an A3S Code `use` worker, call the available + `mcp__use_science__*` tools directly. The host owns installation and MCP + lifecycle; do not run installation or shell commands there. +- In a CLI-only agent host, use `a3s use science ...` commands. + +Select the narrowest authoritative source: + +- Use PubMed for peer-reviewed biomedical literature and article metadata. +- Use bioRxiv for preprints; always label results as preprints. +- Use ChEMBL for molecules, targets, and bioactivity records. +- Use ClinicalTrials.gov for registered study protocols and recruitment status. +- Use Ensembl for gene coordinates, identifiers, and orthologs. + +Start with `science_doctor`. PubMed calls require +`A3S_SCIENCE_CONTACT_EMAIL`; `NCBI_API_KEY` is optional. Other sources do not +require those variables. + +Preserve PMID, DOI, ChEMBL, NCT, and Ensembl identifiers in the answer. State +which source supports each claim, distinguish database metadata from research +conclusions, and report empty or partial results plainly. Never invent missing +records, silently treat a preprint as peer reviewed, or present retrieved data +as diagnosis or medical advice. Cross-check important claims in more than one +source when the task warrants it. + +CLI examples: + +```bash +a3s use science doctor --json +a3s use science pubmed search "CRISPR off-target effects" --limit 10 --json +a3s use science pubmed get 39712345 --json +a3s use science chembl search-molecules aspirin --limit 10 --json +a3s use science chembl activities --molecule CHEMBL25 --limit 20 --json +a3s use science clinical-trials search melanoma --status RECRUITING --json +a3s use science biorxiv search --from 2026-01-01 --to 2026-01-31 --query protein --json +a3s use science ensembl lookup homo_sapiens TP53 --json +a3s use science ensembl homologs homo_sapiens TP53 --target-species mus_musculus --json +``` diff --git a/crates/science/scripts/package.sh b/crates/science/scripts/package.sh new file mode 100755 index 00000000..ce53f664 --- /dev/null +++ b/crates/science/scripts/package.sh @@ -0,0 +1,25 @@ +#!/usr/bin/env bash +set -euo pipefail + +script_dir="$(cd "$(dirname "${BASH_SOURCE[0]}")" && pwd)" +crate_dir="$(cd "${script_dir}/.." && pwd)" +workspace_dir="$(cd "${crate_dir}/../.." && pwd)" +output_dir="${1:-${crate_dir}/dist/a3s-use-science}" + +if [[ -e "${output_dir}" ]]; then + echo "refusing to overwrite existing output: ${output_dir}" >&2 + exit 2 +fi + +cargo build --manifest-path "${workspace_dir}/Cargo.toml" --release --locked -p a3s-use-science + +target_dir="${CARGO_TARGET_DIR:-${workspace_dir}/target}" +mkdir -p "${output_dir}/bin" "${output_dir}/skills/a3s-use-science" +install -m 0755 "${target_dir}/release/a3s-use-science" "${output_dir}/bin/a3s-use-science" +install -m 0644 "${crate_dir}/package/a3s-use-extension.acl" "${output_dir}/a3s-use-extension.acl" +install -m 0644 "${crate_dir}/package/skills/a3s-use-science/SKILL.md" "${output_dir}/skills/a3s-use-science/SKILL.md" +install -m 0644 "${workspace_dir}/LICENSE" "${output_dir}/LICENSE" +install -m 0644 "${crate_dir}/DATA_SOURCES.md" "${output_dir}/DATA_SOURCES.md" +install -m 0644 "${crate_dir}/UPSTREAM.md" "${output_dir}/UPSTREAM.md" + +echo "packaged a3s/science at ${output_dir}" diff --git a/crates/science/src/biorxiv.rs b/crates/science/src/biorxiv.rs new file mode 100644 index 00000000..52088dd4 --- /dev/null +++ b/crates/science/src/biorxiv.rs @@ -0,0 +1,293 @@ +use std::collections::HashSet; +use std::time::Duration; + +use a3s_use_core::{UseError, UseResult}; +use serde::Deserialize; +use serde_json::Value; + +use crate::models::{BioRxivPage, BioRxivRecord}; +use crate::ScienceClient; + +const BIORXIV_INTERVAL: Duration = Duration::from_millis(200); +const MAX_SCAN_RECORDS: usize = 500; + +#[derive(Debug, Deserialize)] +struct BioRxivEnvelope { + #[serde(default)] + messages: Vec, + #[serde(default)] + collection: Vec, +} + +#[derive(Debug, Default, Deserialize)] +struct BioRxivMessage { + #[serde(default)] + total: Value, + #[serde(default)] + count: Value, +} + +#[derive(Debug, Deserialize)] +struct RawBioRxivRecord { + doi: String, + title: String, + authors: String, + #[serde(default, rename = "abstract")] + abstract_text: Option, + #[serde(default)] + category: Option, + #[serde(default)] + date: Option, + #[serde(default)] + version: Option, + #[serde(default, rename = "published")] + published_doi: Option, +} + +impl ScienceClient { + pub async fn biorxiv_search( + &self, + from_date: &str, + to_date: &str, + query: Option<&str>, + category: Option<&str>, + limit: usize, + ) -> UseResult { + validate_date_range(from_date, to_date)?; + let limit = bounded_limit(limit)?; + let query = query + .map(str::trim) + .filter(|query| !query.is_empty()) + .map(str::to_lowercase); + let category = category + .map(str::trim) + .filter(|category| !category.is_empty()) + .map(str::to_lowercase); + + let mut cursor = 0_usize; + let mut scanned = 0_usize; + let mut total_upstream = None; + let mut items = Vec::new(); + let mut seen = HashSet::new(); + while scanned < MAX_SCAN_RECORDS && items.len() < limit { + let cursor_text = cursor.to_string(); + let url = self.endpoint_url( + &self.endpoints.biorxiv, + &[ + "details", + "biorxiv", + from_date, + to_date, + &cursor_text, + "json", + ], + )?; + let envelope: BioRxivEnvelope = self + .get_json("bioRxiv", self.http.get(url), BIORXIV_INTERVAL) + .await?; + let message = envelope.messages.first(); + total_upstream = + total_upstream.or_else(|| message.and_then(|message| flexible_u64(&message.total))); + let reported_count = message + .and_then(|message| flexible_u64(&message.count)) + .map(|count| count as usize) + .unwrap_or(envelope.collection.len()); + let received = envelope.collection.len(); + if received == 0 { + break; + } + scanned = scanned.saturating_add(received); + cursor = cursor.saturating_add(reported_count.max(received)); + for record in envelope.collection { + if !matches_filters(&record, query.as_deref(), category.as_deref()) { + continue; + } + let key = format!( + "{}#{}", + record.doi, + record.version.as_deref().unwrap_or_default() + ); + if seen.insert(key) { + items.push(convert_record(record)); + if items.len() == limit { + break; + } + } + } + if reported_count == 0 || total_upstream.is_some_and(|total| cursor as u64 >= total) { + break; + } + } + Ok(BioRxivPage { + total_upstream, + scanned, + items, + }) + } + + pub async fn biorxiv_get(&self, doi: &str) -> UseResult> { + let suffix = validate_biorxiv_doi(doi)?; + let url = self.endpoint_url( + &self.endpoints.biorxiv, + &["details", "biorxiv", "10.1101", suffix, "na", "json"], + )?; + let envelope: BioRxivEnvelope = self + .get_json("bioRxiv", self.http.get(url), BIORXIV_INTERVAL) + .await?; + if envelope.collection.is_empty() { + return Err(UseError::new( + "use.science.not_found", + format!("bioRxiv did not return DOI {doi}."), + ) + .with_detail("service", "bioRxiv") + .with_detail("doi", doi)); + } + Ok(envelope + .collection + .into_iter() + .map(convert_record) + .collect()) + } +} + +fn convert_record(record: RawBioRxivRecord) -> BioRxivRecord { + BioRxivRecord { + doi: record.doi, + title: record.title, + authors: record.authors, + abstract_text: non_empty(record.abstract_text), + category: non_empty(record.category), + date: non_empty(record.date), + version: non_empty(record.version), + published_doi: non_empty(record.published_doi).filter(|value| value != "NA"), + } +} + +fn non_empty(value: Option) -> Option { + value.filter(|value| !value.trim().is_empty()) +} + +fn matches_filters(record: &RawBioRxivRecord, query: Option<&str>, category: Option<&str>) -> bool { + let query_matches = query.is_none_or(|query| { + [ + record.title.as_str(), + record.authors.as_str(), + record.abstract_text.as_deref().unwrap_or_default(), + record.doi.as_str(), + ] + .iter() + .any(|value| value.to_lowercase().contains(query)) + }); + let category_matches = category.is_none_or(|category| { + record + .category + .as_deref() + .is_some_and(|value| value.eq_ignore_ascii_case(category)) + }); + query_matches && category_matches +} + +fn flexible_u64(value: &Value) -> Option { + match value { + Value::Number(value) => value.as_u64(), + Value::String(value) => value.parse().ok(), + _ => None, + } +} + +fn validate_date_range(from_date: &str, to_date: &str) -> UseResult<()> { + if !valid_iso_date(from_date) || !valid_iso_date(to_date) || from_date > to_date { + return Err(UseError::new( + "use.science.date_invalid", + "bioRxiv dates must form an ordered YYYY-MM-DD range.", + )); + } + Ok(()) +} + +fn valid_iso_date(value: &str) -> bool { + let parts = value.split('-').collect::>(); + let (Ok(year), Ok(month), Ok(day)) = ( + parts.first().unwrap_or(&"").parse::(), + parts.get(1).unwrap_or(&"").parse::(), + parts.get(2).unwrap_or(&"").parse::(), + ) else { + return false; + }; + if parts.len() != 3 || parts[0].len() != 4 || parts[1].len() != 2 || parts[2].len() != 2 { + return false; + } + let leap = year % 4 == 0 && (year % 100 != 0 || year % 400 == 0); + let max_day = match month { + 1 | 3 | 5 | 7 | 8 | 10 | 12 => 31, + 4 | 6 | 9 | 11 => 30, + 2 if leap => 29, + 2 => 28, + _ => return false, + }; + (1900..=9999).contains(&year) && (1..=max_day).contains(&day) +} + +fn validate_biorxiv_doi(doi: &str) -> UseResult<&str> { + let suffix = doi.strip_prefix("10.1101/").unwrap_or_default(); + if suffix.is_empty() + || suffix.len() > 200 + || !suffix.bytes().all(|byte| { + byte.is_ascii_alphanumeric() || matches!(byte, b'.' | b'-' | b'_' | b'(' | b')') + }) + { + return Err(UseError::new( + "use.science.identifier_invalid", + "A bioRxiv DOI must start with 10.1101/ and contain a safe DOI suffix.", + )); + } + Ok(suffix) +} + +fn bounded_limit(limit: usize) -> UseResult { + if !(1..=100).contains(&limit) { + return Err(UseError::new( + "use.science.limit_invalid", + "bioRxiv result limit must be between 1 and 100.", + )); + } + Ok(limit) +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn validates_real_calendar_dates_and_dois() { + assert!(valid_iso_date("2024-02-29")); + assert!(!valid_iso_date("2025-02-29")); + assert!(validate_biorxiv_doi("10.1101/2026.01.01.123456").is_ok()); + assert_eq!( + validate_biorxiv_doi("https://example.com") + .unwrap_err() + .code, + "use.science.identifier_invalid" + ); + } + + #[test] + fn filters_records_without_losing_case_insensitivity() { + let record = RawBioRxivRecord { + doi: "10.1101/example".to_string(), + title: "Protein Design".to_string(), + authors: "A. Author".to_string(), + abstract_text: Some("A diffusion model".to_string()), + category: Some("Bioinformatics".to_string()), + date: None, + version: None, + published_doi: None, + }; + assert!(matches_filters( + &record, + Some("protein"), + Some("bioinformatics") + )); + assert!(!matches_filters(&record, Some("genome"), None)); + } +} diff --git a/crates/science/src/chembl.rs b/crates/science/src/chembl.rs new file mode 100644 index 00000000..5043f570 --- /dev/null +++ b/crates/science/src/chembl.rs @@ -0,0 +1,270 @@ +use std::time::Duration; + +use a3s_use_core::{UseError, UseResult}; +use serde::Deserialize; +use serde_json::Value; + +use crate::models::{ChemblActivity, ChemblMolecule, ChemblTarget, Page}; +use crate::ScienceClient; + +const CHEMBL_INTERVAL: Duration = Duration::from_millis(100); + +#[derive(Debug, Default, Deserialize)] +struct PageMeta { + #[serde(default)] + total_count: Option, + #[serde(default)] + next: Option, +} + +#[derive(Debug, Deserialize)] +struct MoleculeEnvelope { + #[serde(default)] + molecules: Vec, + #[serde(default)] + page_meta: PageMeta, +} + +#[derive(Debug, Deserialize)] +struct TargetEnvelope { + #[serde(default)] + targets: Vec, + #[serde(default)] + page_meta: PageMeta, +} + +#[derive(Debug, Deserialize)] +struct ActivityEnvelope { + #[serde(default)] + activities: Vec, + #[serde(default)] + page_meta: PageMeta, +} + +impl ScienceClient { + pub async fn chembl_search_molecules( + &self, + query: &str, + limit: usize, + ) -> UseResult> { + let query = required_query(query)?; + let limit = bounded_limit(limit)?; + let url = self.endpoint_url(&self.endpoints.chembl, &["molecule", "search.json"])?; + let params = [("q", query.to_string()), ("limit", limit.to_string())]; + let envelope: MoleculeEnvelope = self + .get_json("ChEMBL", self.http.get(url).query(¶ms), CHEMBL_INTERVAL) + .await?; + Ok(Page { + total: envelope.page_meta.total_count, + next_page_token: envelope.page_meta.next, + items: envelope + .molecules + .iter() + .filter_map(parse_molecule) + .collect(), + }) + } + + pub async fn chembl_get_molecule(&self, chembl_id: &str) -> UseResult { + validate_chembl_id(chembl_id)?; + let url = self.endpoint_url( + &self.endpoints.chembl, + &["molecule", &format!("{chembl_id}.json")], + )?; + let value: Value = self + .get_json("ChEMBL", self.http.get(url), CHEMBL_INTERVAL) + .await?; + parse_molecule(&value).ok_or_else(|| { + UseError::new( + "use.science.response_invalid", + "ChEMBL returned a molecule without a molecule_chembl_id.", + ) + }) + } + + pub async fn chembl_search_targets( + &self, + query: &str, + limit: usize, + ) -> UseResult> { + let query = required_query(query)?; + let limit = bounded_limit(limit)?; + let url = self.endpoint_url(&self.endpoints.chembl, &["target", "search.json"])?; + let params = [("q", query.to_string()), ("limit", limit.to_string())]; + let envelope: TargetEnvelope = self + .get_json("ChEMBL", self.http.get(url).query(¶ms), CHEMBL_INTERVAL) + .await?; + Ok(Page { + total: envelope.page_meta.total_count, + next_page_token: envelope.page_meta.next, + items: envelope.targets.iter().filter_map(parse_target).collect(), + }) + } + + pub async fn chembl_activities( + &self, + molecule_chembl_id: Option<&str>, + target_chembl_id: Option<&str>, + limit: usize, + ) -> UseResult> { + let limit = bounded_limit(limit)?; + if molecule_chembl_id.is_none() && target_chembl_id.is_none() { + return Err(UseError::new( + "use.science.input_invalid", + "ChEMBL activities require a molecule or target ChEMBL ID.", + )); + } + if let Some(identifier) = molecule_chembl_id { + validate_chembl_id(identifier)?; + } + if let Some(identifier) = target_chembl_id { + validate_chembl_id(identifier)?; + } + let url = self.endpoint_url(&self.endpoints.chembl, &["activity.json"])?; + let mut query = vec![("limit", limit.to_string())]; + if let Some(identifier) = molecule_chembl_id { + query.push(("molecule_chembl_id", identifier.to_string())); + } + if let Some(identifier) = target_chembl_id { + query.push(("target_chembl_id", identifier.to_string())); + } + let envelope: ActivityEnvelope = self + .get_json("ChEMBL", self.http.get(url).query(&query), CHEMBL_INTERVAL) + .await?; + Ok(Page { + total: envelope.page_meta.total_count, + next_page_token: envelope.page_meta.next, + items: envelope.activities.iter().map(parse_activity).collect(), + }) + } +} + +fn parse_molecule(value: &Value) -> Option { + Some(ChemblMolecule { + chembl_id: value.get("molecule_chembl_id")?.as_str()?.to_string(), + preferred_name: value_string(value.get("pref_name")), + molecule_type: value_string(value.get("molecule_type")), + max_phase: value.get("max_phase").and_then(value_f64), + canonical_smiles: value + .get("molecule_structures") + .and_then(|structures| structures.get("canonical_smiles")) + .and_then(Value::as_str) + .map(str::to_string), + standard_inchi_key: value + .get("molecule_structures") + .and_then(|structures| structures.get("standard_inchi_key")) + .and_then(Value::as_str) + .map(str::to_string), + }) +} + +fn parse_target(value: &Value) -> Option { + Some(ChemblTarget { + chembl_id: value.get("target_chembl_id")?.as_str()?.to_string(), + preferred_name: value_string(value.get("pref_name")), + target_type: value_string(value.get("target_type")), + organism: value_string(value.get("organism")), + }) +} + +fn parse_activity(value: &Value) -> ChemblActivity { + ChemblActivity { + activity_id: value_string(value.get("activity_id")), + molecule_chembl_id: value_string(value.get("molecule_chembl_id")), + target_chembl_id: value_string(value.get("target_chembl_id")), + assay_chembl_id: value_string(value.get("assay_chembl_id")), + standard_type: value_string(value.get("standard_type")), + standard_relation: value_string(value.get("standard_relation")), + standard_value: value_string(value.get("standard_value")), + standard_units: value_string(value.get("standard_units")), + pchembl_value: value_string(value.get("pchembl_value")), + } +} + +fn value_string(value: Option<&Value>) -> Option { + match value? { + Value::String(value) if !value.is_empty() => Some(value.clone()), + Value::Number(value) => Some(value.to_string()), + _ => None, + } +} + +fn value_f64(value: &Value) -> Option { + match value { + Value::Number(value) => value.as_f64(), + Value::String(value) => value.parse().ok(), + _ => None, + } +} + +fn required_query(query: &str) -> UseResult<&str> { + let query = query.trim(); + if query.is_empty() { + return Err(UseError::new( + "use.science.input_invalid", + "ChEMBL query cannot be empty.", + )); + } + Ok(query) +} + +fn bounded_limit(limit: usize) -> UseResult { + if !(1..=100).contains(&limit) { + return Err(UseError::new( + "use.science.limit_invalid", + "ChEMBL result limit must be between 1 and 100.", + )); + } + Ok(limit) +} + +fn validate_chembl_id(identifier: &str) -> UseResult<()> { + let suffix = identifier.strip_prefix("CHEMBL").unwrap_or_default(); + if suffix.is_empty() || !suffix.bytes().all(|byte| byte.is_ascii_digit()) { + return Err(UseError::new( + "use.science.identifier_invalid", + "A ChEMBL identifier must use the form CHEMBL followed by digits.", + )); + } + Ok(()) +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn parses_typed_chembl_records() { + let molecule = serde_json::json!({ + "molecule_chembl_id": "CHEMBL25", + "pref_name": "ASPIRIN", + "molecule_type": "Small molecule", + "max_phase": 4, + "molecule_structures": { + "canonical_smiles": "CC(=O)OC1=CC=CC=C1C(=O)O", + "standard_inchi_key": "BSYNRYMUTXBXSQ-UHFFFAOYSA-N" + } + }); + let parsed = parse_molecule(&molecule).unwrap(); + assert_eq!(parsed.chembl_id, "CHEMBL25"); + assert_eq!(parsed.max_phase, Some(4.0)); + + let activity = parse_activity(&serde_json::json!({ + "activity_id": 42, + "standard_value": "12.5" + })); + assert_eq!(activity.activity_id.as_deref(), Some("42")); + } + + #[test] + fn rejects_unbounded_or_malformed_inputs() { + assert_eq!( + validate_chembl_id("../CHEMBL25").unwrap_err().code, + "use.science.identifier_invalid" + ); + assert_eq!( + bounded_limit(101).unwrap_err().code, + "use.science.limit_invalid" + ); + } +} diff --git a/crates/science/src/cli.rs b/crates/science/src/cli.rs new file mode 100644 index 00000000..5aaa1e82 --- /dev/null +++ b/crates/science/src/cli.rs @@ -0,0 +1,370 @@ +use a3s_use_core::{UseError, UseResult}; +use clap::error::ErrorKind; +use clap::{Args, Parser, Subcommand}; +use serde::Serialize; + +use crate::{ScienceClient, ScienceMcpServer}; + +#[derive(Debug)] +pub struct CommandOutput { + pub human: String, + pub json: serde_json::Value, + pub exit_code: u8, + pub should_print: bool, +} + +impl CommandOutput { + fn data(value: T) -> UseResult + where + T: Serialize, + { + let data = serde_json::to_value(value).map_err(output_error)?; + let human = serde_json::to_string_pretty(&data).map_err(output_error)?; + Ok(Self { + human, + json: serde_json::json!({ + "schemaVersion": 1, + "ok": true, + "data": data, + }), + exit_code: 0, + should_print: true, + }) + } + + fn text(value: String) -> Self { + Self { + human: value.clone(), + json: serde_json::json!({ + "schemaVersion": 1, + "ok": true, + "data": { "text": value }, + }), + exit_code: 0, + should_print: true, + } + } + + fn silent() -> Self { + Self { + human: String::new(), + json: serde_json::Value::Null, + exit_code: 0, + should_print: false, + } + } +} + +#[derive(Debug, Parser)] +#[command( + name = "a3s-use-science", + version, + about = "Read-only life-science data tools for A3S Use", + arg_required_else_help = true +)] +struct Cli { + /// Emit one versioned JSON document. + #[arg(long, global = true)] + json: bool, + + #[command(subcommand)] + command: Command, +} + +#[derive(Debug, Subcommand)] +enum Command { + /// Inspect local configuration without making a network request. + Doctor, + /// Search or retrieve PubMed article summaries. + Pubmed(PubmedArgs), + /// Search ChEMBL molecules, targets, and activities. + Chembl(ChemblArgs), + /// Search or retrieve ClinicalTrials.gov studies. + #[command(name = "clinical-trials")] + ClinicalTrials(ClinicalTrialsArgs), + /// Search or retrieve bioRxiv preprints. + Biorxiv(BioRxivArgs), + /// Look up Ensembl genes and homologs. + Ensembl(EnsemblArgs), + /// Run an extension protocol surface. + Serve(ServeArgs), +} + +#[derive(Debug, Args)] +struct PubmedArgs { + #[command(subcommand)] + command: PubmedCommand, +} + +#[derive(Debug, Subcommand)] +enum PubmedCommand { + /// Search PubMed and return article summaries. + Search { + query: String, + #[arg(long, default_value_t = 20)] + limit: usize, + }, + /// Retrieve one PubMed article summary by PMID. + Get { pmid: String }, +} + +#[derive(Debug, Args)] +struct ChemblArgs { + #[command(subcommand)] + command: ChemblCommand, +} + +#[derive(Debug, Subcommand)] +enum ChemblCommand { + /// Search ChEMBL molecules. + #[command(name = "search-molecules")] + SearchMolecules { + query: String, + #[arg(long, default_value_t = 20)] + limit: usize, + }, + /// Retrieve one ChEMBL molecule. + #[command(name = "get-molecule")] + GetMolecule { chembl_id: String }, + /// Search ChEMBL targets. + #[command(name = "search-targets")] + SearchTargets { + query: String, + #[arg(long, default_value_t = 20)] + limit: usize, + }, + /// Retrieve bioactivity records for a molecule, target, or both. + Activities { + #[arg(long = "molecule")] + molecule_chembl_id: Option, + #[arg(long = "target")] + target_chembl_id: Option, + #[arg(long, default_value_t = 20)] + limit: usize, + }, +} + +#[derive(Debug, Args)] +struct ClinicalTrialsArgs { + #[command(subcommand)] + command: ClinicalTrialsCommand, +} + +#[derive(Debug, Subcommand)] +enum ClinicalTrialsCommand { + /// Search ClinicalTrials.gov studies. + Search { + query: String, + #[arg(long = "status")] + statuses: Vec, + #[arg(long, default_value_t = 20)] + limit: usize, + #[arg(long)] + page_token: Option, + }, + /// Retrieve one study by NCT identifier. + Get { nct_id: String }, +} + +#[derive(Debug, Args)] +struct BioRxivArgs { + #[command(subcommand)] + command: BioRxivCommand, +} + +#[derive(Debug, Subcommand)] +enum BioRxivCommand { + /// Search a bounded bioRxiv date range. + Search { + #[arg(long = "from")] + from_date: String, + #[arg(long = "to")] + to_date: String, + #[arg(long)] + query: Option, + #[arg(long)] + category: Option, + #[arg(long, default_value_t = 20)] + limit: usize, + }, + /// Retrieve all returned versions of one bioRxiv DOI. + Get { doi: String }, +} + +#[derive(Debug, Args)] +struct EnsemblArgs { + #[command(subcommand)] + command: EnsemblCommand, +} + +#[derive(Debug, Subcommand)] +enum EnsemblCommand { + /// Look up one gene by species and symbol. + Lookup { species: String, symbol: String }, + /// Retrieve orthologs for one gene symbol. + Homologs { + species: String, + symbol: String, + #[arg(long)] + target_species: Option, + #[arg(long, default_value_t = 50)] + limit: usize, + }, +} + +#[derive(Debug, Args)] +struct ServeArgs { + /// Serve standard MCP over stdin/stdout. + #[arg(long)] + mcp: bool, +} + +pub async fn run(args: Vec) -> UseResult { + let mut argv = vec!["a3s-use-science".to_string()]; + argv.extend(args); + let cli = match Cli::try_parse_from(argv) { + Ok(cli) => cli, + Err(error) + if matches!( + error.kind(), + ErrorKind::DisplayHelp | ErrorKind::DisplayVersion + ) => + { + return Ok(CommandOutput::text(error.to_string())); + } + Err(error) => return Err(usage_error(error.to_string())), + }; + + if let Command::Serve(serve) = &cli.command { + if !serve.mcp { + return Err(usage_error("serve requires --mcp")); + } + if cli.json { + return Err(usage_error("--json cannot be combined with serve --mcp")); + } + ScienceMcpServer::from_env()?.serve_stdio().await?; + return Ok(CommandOutput::silent()); + } + + let client = ScienceClient::from_env()?; + match cli.command { + Command::Doctor => CommandOutput::data(client.diagnostic()), + Command::Pubmed(args) => match args.command { + PubmedCommand::Search { query, limit } => { + CommandOutput::data(client.pubmed_search(&query, limit).await?) + } + PubmedCommand::Get { pmid } => CommandOutput::data(client.pubmed_get(&pmid).await?), + }, + Command::Chembl(args) => match args.command { + ChemblCommand::SearchMolecules { query, limit } => { + CommandOutput::data(client.chembl_search_molecules(&query, limit).await?) + } + ChemblCommand::GetMolecule { chembl_id } => { + CommandOutput::data(client.chembl_get_molecule(&chembl_id).await?) + } + ChemblCommand::SearchTargets { query, limit } => { + CommandOutput::data(client.chembl_search_targets(&query, limit).await?) + } + ChemblCommand::Activities { + molecule_chembl_id, + target_chembl_id, + limit, + } => CommandOutput::data( + client + .chembl_activities( + molecule_chembl_id.as_deref(), + target_chembl_id.as_deref(), + limit, + ) + .await?, + ), + }, + Command::ClinicalTrials(args) => match args.command { + ClinicalTrialsCommand::Search { + query, + statuses, + limit, + page_token, + } => CommandOutput::data( + client + .clinical_trials_search(&query, &statuses, limit, page_token.as_deref()) + .await?, + ), + ClinicalTrialsCommand::Get { nct_id } => { + CommandOutput::data(client.clinical_trial_get(&nct_id).await?) + } + }, + Command::Biorxiv(args) => match args.command { + BioRxivCommand::Search { + from_date, + to_date, + query, + category, + limit, + } => CommandOutput::data( + client + .biorxiv_search( + &from_date, + &to_date, + query.as_deref(), + category.as_deref(), + limit, + ) + .await?, + ), + BioRxivCommand::Get { doi } => CommandOutput::data(client.biorxiv_get(&doi).await?), + }, + Command::Ensembl(args) => match args.command { + EnsemblCommand::Lookup { species, symbol } => { + CommandOutput::data(client.ensembl_lookup_gene(&species, &symbol).await?) + } + EnsemblCommand::Homologs { + species, + symbol, + target_species, + limit, + } => CommandOutput::data( + client + .ensembl_homologs(&species, &symbol, target_species.as_deref(), limit) + .await?, + ), + }, + Command::Serve(_) => Err(UseError::new( + "use.science.command_invalid", + "Science MCP command dispatch reached an invalid state.", + )), + } +} + +fn output_error(error: serde_json::Error) -> UseError { + UseError::new( + "use.science.output_invalid", + format!("Failed to encode science command output: {error}"), + ) +} + +fn usage_error(message: impl Into) -> UseError { + UseError::new("use.science.usage_invalid", message) + .with_suggestion("Run 'a3s use science --help'.") +} + +#[cfg(test)] +mod tests { + use super::*; + + #[tokio::test] + async fn doctor_is_versioned_and_does_not_require_network_configuration() { + let output = run(vec!["doctor".to_string(), "--json".to_string()]) + .await + .unwrap(); + assert_eq!(output.json["schemaVersion"], 1); + assert_eq!(output.json["ok"], true); + assert_eq!(output.json["data"]["sources"].as_array().unwrap().len(), 5); + } + + #[tokio::test] + async fn serve_requires_an_explicit_protocol() { + let error = run(vec!["serve".to_string()]).await.unwrap_err(); + assert_eq!(error.code, "use.science.usage_invalid"); + } +} diff --git a/crates/science/src/client.rs b/crates/science/src/client.rs new file mode 100644 index 00000000..7d030858 --- /dev/null +++ b/crates/science/src/client.rs @@ -0,0 +1,365 @@ +use std::collections::HashMap; +use std::sync::Arc; +use std::time::Duration; + +use a3s_use_core::{UseError, UseResult}; +use reqwest::{RequestBuilder, Response}; +use serde::de::DeserializeOwned; +use tokio::sync::Mutex; +use tokio::time::Instant; +use url::Url; + +use crate::models::ScienceDiagnostic; + +const DEFAULT_TIMEOUT: Duration = Duration::from_secs(30); +const ERROR_BODY_LIMIT: usize = 1_024; + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct ScienceEndpoints { + pub pubmed: Url, + pub chembl: Url, + pub clinical_trials: Url, + pub biorxiv: Url, + pub ensembl: Url, +} + +impl ScienceEndpoints { + /// Return the public upstream endpoints used by the toolkit. + pub fn public() -> UseResult { + Ok(Self { + pubmed: parse_endpoint("PubMed", "https://eutils.ncbi.nlm.nih.gov/entrez/eutils/")?, + chembl: parse_endpoint("ChEMBL", "https://www.ebi.ac.uk/chembl/api/data/")?, + clinical_trials: parse_endpoint( + "ClinicalTrials.gov", + "https://clinicaltrials.gov/api/v2/", + )?, + biorxiv: parse_endpoint("bioRxiv", "https://api.biorxiv.org/")?, + ensembl: parse_endpoint("Ensembl", "https://rest.ensembl.org/")?, + }) + } +} + +#[derive(Debug, Clone)] +pub struct ScienceClientBuilder { + endpoints: Option, + contact_email: Option, + ncbi_api_key: Option, + timeout: Duration, + user_agent: Option, +} + +impl Default for ScienceClientBuilder { + fn default() -> Self { + Self { + endpoints: None, + contact_email: None, + ncbi_api_key: None, + timeout: DEFAULT_TIMEOUT, + user_agent: None, + } + } +} + +impl ScienceClientBuilder { + pub fn endpoints(mut self, endpoints: ScienceEndpoints) -> Self { + self.endpoints = Some(endpoints); + self + } + + pub fn contact_email(mut self, contact_email: impl Into) -> Self { + self.contact_email = Some(contact_email.into()); + self + } + + pub fn ncbi_api_key(mut self, api_key: impl Into) -> Self { + self.ncbi_api_key = Some(api_key.into()); + self + } + + pub fn timeout(mut self, timeout: Duration) -> Self { + self.timeout = timeout; + self + } + + pub fn user_agent(mut self, user_agent: impl Into) -> Self { + self.user_agent = Some(user_agent.into()); + self + } + + pub fn build(self) -> UseResult { + if self.timeout.is_zero() { + return Err(UseError::new( + "use.science.config_invalid", + "Science request timeout must be greater than zero.", + )); + } + if let Some(email) = self.contact_email.as_deref() { + if !email.contains('@') || email.chars().any(char::is_whitespace) { + return Err(UseError::new( + "use.science.contact_email_invalid", + "A3S_SCIENCE_CONTACT_EMAIL must contain a valid contact email address.", + )); + } + } + let endpoints = match self.endpoints { + Some(endpoints) => endpoints, + None => ScienceEndpoints::public()?, + }; + for (name, endpoint) in endpoint_pairs(&endpoints) { + if !matches!(endpoint.scheme(), "http" | "https") { + return Err(UseError::new( + "use.science.config_invalid", + format!("The {name} endpoint must use HTTP or HTTPS."), + )); + } + } + + let user_agent = self + .user_agent + .unwrap_or_else(|| format!("a3s-use-science/{}", env!("CARGO_PKG_VERSION"))); + let http = reqwest::Client::builder() + .timeout(self.timeout) + .user_agent(user_agent) + .build() + .map_err(|error| { + UseError::new( + "use.science.client_invalid", + format!("Failed to construct the science HTTP client: {error}"), + ) + })?; + + Ok(ScienceClient { + http, + endpoints, + contact_email: self.contact_email, + ncbi_api_key: self.ncbi_api_key, + gate: Arc::new(RequestGate::default()), + }) + } +} + +#[derive(Debug, Clone)] +pub struct ScienceClient { + pub(crate) http: reqwest::Client, + pub(crate) endpoints: ScienceEndpoints, + pub(crate) contact_email: Option, + pub(crate) ncbi_api_key: Option, + gate: Arc, +} + +impl ScienceClient { + pub fn builder() -> ScienceClientBuilder { + ScienceClientBuilder::default() + } + + pub fn from_env() -> UseResult { + let mut builder = Self::builder(); + if let Some(contact_email) = optional_env("A3S_SCIENCE_CONTACT_EMAIL")? { + builder = builder.contact_email(contact_email); + } + if let Some(api_key) = optional_env("NCBI_API_KEY")? { + builder = builder.ncbi_api_key(api_key); + } + builder.build() + } + + pub fn diagnostic(&self) -> ScienceDiagnostic { + ScienceDiagnostic { + version: env!("CARGO_PKG_VERSION").to_string(), + contact_email_configured: self.contact_email.is_some(), + ncbi_api_key_configured: self.ncbi_api_key.is_some(), + sources: vec![ + "PubMed".to_string(), + "ChEMBL".to_string(), + "ClinicalTrials.gov".to_string(), + "bioRxiv".to_string(), + "Ensembl".to_string(), + ], + message: "Read-only public life-science data sources are configured; no network request was made." + .to_string(), + } + } + + pub(crate) async fn get_json( + &self, + service: &'static str, + request: RequestBuilder, + min_interval: Duration, + ) -> UseResult + where + T: DeserializeOwned, + { + self.gate.wait(service, min_interval).await; + let response = request.send().await.map_err(|error| { + UseError::new( + "use.science.upstream_unavailable", + format!("{service} request failed: {error}"), + ) + .with_detail("service", service) + })?; + parse_json(service, response).await + } + + pub(crate) fn endpoint_url(&self, base: &Url, segments: &[&str]) -> UseResult { + let mut url = base.clone(); + { + let mut path = url.path_segments_mut().map_err(|_| { + UseError::new( + "use.science.config_invalid", + format!("Endpoint '{base}' cannot be used as a hierarchical URL."), + ) + })?; + path.pop_if_empty(); + path.extend(segments.iter().copied()); + } + Ok(url) + } +} + +async fn parse_json(service: &'static str, response: Response) -> UseResult +where + T: DeserializeOwned, +{ + let status = response.status(); + if !status.is_success() { + let body = bounded_response_text(response, ERROR_BODY_LIMIT).await; + return Err(UseError::new( + "use.science.upstream_error", + format!("{service} returned HTTP {status}."), + ) + .with_detail("service", service) + .with_detail("status", u64::from(status.as_u16())) + .with_detail("body", body)); + } + response.json::().await.map_err(|error| { + UseError::new( + "use.science.response_invalid", + format!("{service} returned an invalid JSON response: {error}"), + ) + .with_detail("service", service) + }) +} + +async fn bounded_response_text(mut response: Response, max_bytes: usize) -> String { + let mut bytes = Vec::with_capacity(max_bytes.saturating_add(1)); + while bytes.len() <= max_bytes { + let chunk = match response.chunk().await { + Ok(Some(chunk)) if !chunk.is_empty() => chunk, + Ok(Some(_)) | Ok(None) | Err(_) => break, + }; + let remaining = max_bytes.saturating_add(1).saturating_sub(bytes.len()); + let take = remaining.min(chunk.len()); + bytes.extend_from_slice(&chunk[..take]); + if take < chunk.len() { + break; + } + } + let truncated = bytes.len() > max_bytes; + bytes.truncate(max_bytes); + let mut output = String::from_utf8_lossy(&bytes).into_owned(); + if truncated { + output.push('…'); + } + output +} + +fn endpoint_pairs(endpoints: &ScienceEndpoints) -> [(&'static str, &Url); 5] { + [ + ("PubMed", &endpoints.pubmed), + ("ChEMBL", &endpoints.chembl), + ("ClinicalTrials.gov", &endpoints.clinical_trials), + ("bioRxiv", &endpoints.biorxiv), + ("Ensembl", &endpoints.ensembl), + ] +} + +fn parse_endpoint(service: &'static str, value: &'static str) -> UseResult { + Url::parse(value).map_err(|error| { + UseError::new( + "use.science.config_invalid", + format!("The built-in {service} endpoint is invalid: {error}"), + ) + }) +} + +fn optional_env(name: &'static str) -> UseResult> { + match std::env::var(name) { + Ok(value) => Ok(Some(value)), + Err(std::env::VarError::NotPresent) => Ok(None), + Err(std::env::VarError::NotUnicode(_)) => Err(UseError::new( + "use.science.config_invalid", + format!("Environment variable {name} must contain valid UTF-8."), + )), + } +} + +#[derive(Debug, Default)] +struct RequestGate { + next_allowed: Mutex>, +} + +impl RequestGate { + async fn wait(&self, service: &'static str, min_interval: Duration) { + if min_interval.is_zero() { + return; + } + let now = Instant::now(); + let start = { + let mut next_allowed = self.next_allowed.lock().await; + let start = next_allowed.get(service).copied().unwrap_or(now).max(now); + next_allowed.insert(service, start + min_interval); + start + }; + if start > now { + tokio::time::sleep_until(start).await; + } + } +} + +#[cfg(test)] +mod tests { + use super::*; + + fn assert_send_sync() {} + + #[test] + fn rejects_invalid_contact_email_and_non_http_endpoints() { + assert_eq!( + ScienceClient::builder() + .contact_email("not-an-email") + .build() + .unwrap_err() + .code, + "use.science.contact_email_invalid" + ); + + let mut endpoints = ScienceEndpoints::public().unwrap(); + endpoints.pubmed = Url::parse("file:///tmp/pubmed").unwrap(); + assert_eq!( + ScienceClient::builder() + .endpoints(endpoints) + .build() + .unwrap_err() + .code, + "use.science.config_invalid" + ); + } + + #[test] + fn endpoint_builder_percent_encodes_untrusted_segments() { + let client = ScienceClient::builder().build().unwrap(); + let url = client + .endpoint_url( + &client.endpoints.ensembl, + &["lookup", "symbol", "human", "A/B"], + ) + .unwrap(); + assert!(url.as_str().ends_with("/lookup/symbol/human/A%2FB")); + } + + #[test] + fn public_clients_are_send_and_sync() { + assert_send_sync::(); + assert_send_sync::(); + } +} diff --git a/crates/science/src/clinical_trials.rs b/crates/science/src/clinical_trials.rs new file mode 100644 index 00000000..2dc31ddf --- /dev/null +++ b/crates/science/src/clinical_trials.rs @@ -0,0 +1,312 @@ +use std::time::Duration; + +use a3s_use_core::{UseError, UseResult}; +use serde::Deserialize; + +use crate::models::{ClinicalTrial, ClinicalTrialPage}; +use crate::ScienceClient; + +const CLINICAL_TRIALS_INTERVAL: Duration = Duration::from_millis(120); + +#[derive(Debug, Deserialize)] +#[serde(rename_all = "camelCase")] +struct StudiesEnvelope { + #[serde(default)] + studies: Vec, + #[serde(default)] + next_page_token: Option, + #[serde(default)] + total_count: Option, +} + +#[derive(Debug, Deserialize)] +#[serde(rename_all = "camelCase")] +struct Study { + protocol_section: ProtocolSection, +} + +#[derive(Debug, Deserialize)] +#[serde(rename_all = "camelCase")] +struct ProtocolSection { + identification_module: IdentificationModule, + #[serde(default)] + status_module: StatusModule, + #[serde(default)] + design_module: DesignModule, + #[serde(default)] + conditions_module: ConditionsModule, + #[serde(default)] + arms_interventions_module: ArmsInterventionsModule, + #[serde(default)] + sponsor_collaborators_module: SponsorCollaboratorsModule, +} + +#[derive(Debug, Deserialize)] +#[serde(rename_all = "camelCase")] +struct IdentificationModule { + nct_id: String, + brief_title: String, + #[serde(default)] + official_title: Option, +} + +#[derive(Debug, Default, Deserialize)] +#[serde(rename_all = "camelCase")] +struct StatusModule { + #[serde(default)] + overall_status: Option, + #[serde(default)] + start_date_struct: Option, + #[serde(default)] + completion_date_struct: Option, +} + +#[derive(Debug, Deserialize)] +#[serde(rename_all = "camelCase")] +struct DateStruct { + date: String, +} + +#[derive(Debug, Default, Deserialize)] +#[serde(rename_all = "camelCase")] +struct DesignModule { + #[serde(default)] + study_type: Option, + #[serde(default)] + phases: Vec, + #[serde(default)] + enrollment_info: Option, +} + +#[derive(Debug, Deserialize)] +struct EnrollmentInfo { + #[serde(default)] + count: Option, +} + +#[derive(Debug, Default, Deserialize)] +struct ConditionsModule { + #[serde(default)] + conditions: Vec, +} + +#[derive(Debug, Default, Deserialize)] +struct ArmsInterventionsModule { + #[serde(default)] + interventions: Vec, +} + +#[derive(Debug, Deserialize)] +struct Intervention { + name: String, +} + +#[derive(Debug, Default, Deserialize)] +#[serde(rename_all = "camelCase")] +struct SponsorCollaboratorsModule { + #[serde(default)] + lead_sponsor: Option, +} + +#[derive(Debug, Deserialize)] +struct LeadSponsor { + name: String, +} + +impl ScienceClient { + pub async fn clinical_trials_search( + &self, + query: &str, + statuses: &[String], + limit: usize, + page_token: Option<&str>, + ) -> UseResult { + let query = query.trim(); + if query.is_empty() { + return Err(UseError::new( + "use.science.input_invalid", + "ClinicalTrials.gov query cannot be empty.", + )); + } + let limit = bounded_limit(limit)?; + for status in statuses { + validate_status(status)?; + } + if let Some(token) = page_token { + validate_page_token(token)?; + } + let url = self.endpoint_url(&self.endpoints.clinical_trials, &["studies"])?; + let mut params = vec![ + ("query.term", query.to_string()), + ("pageSize", limit.to_string()), + ("countTotal", "true".to_string()), + ("format", "json".to_string()), + ]; + if !statuses.is_empty() { + params.push(("filter.overallStatus", statuses.join("|"))); + } + if let Some(token) = page_token { + params.push(("pageToken", token.to_string())); + } + let envelope: StudiesEnvelope = self + .get_json( + "ClinicalTrials.gov", + self.http.get(url).query(¶ms), + CLINICAL_TRIALS_INTERVAL, + ) + .await?; + Ok(ClinicalTrialPage { + total: envelope.total_count, + next_page_token: envelope.next_page_token, + items: envelope.studies.into_iter().map(flatten_study).collect(), + }) + } + + pub async fn clinical_trial_get(&self, nct_id: &str) -> UseResult { + validate_nct_id(nct_id)?; + let url = self.endpoint_url(&self.endpoints.clinical_trials, &["studies", nct_id])?; + let study: Study = self + .get_json( + "ClinicalTrials.gov", + self.http.get(url).query(&[("format", "json")]), + CLINICAL_TRIALS_INTERVAL, + ) + .await?; + Ok(flatten_study(study)) + } +} + +fn flatten_study(study: Study) -> ClinicalTrial { + let protocol = study.protocol_section; + ClinicalTrial { + nct_id: protocol.identification_module.nct_id, + brief_title: protocol.identification_module.brief_title, + official_title: protocol.identification_module.official_title, + overall_status: protocol.status_module.overall_status, + study_type: protocol.design_module.study_type, + phases: protocol.design_module.phases, + conditions: protocol.conditions_module.conditions, + interventions: protocol + .arms_interventions_module + .interventions + .into_iter() + .map(|intervention| intervention.name) + .collect(), + lead_sponsor: protocol + .sponsor_collaborators_module + .lead_sponsor + .map(|sponsor| sponsor.name), + enrollment: protocol + .design_module + .enrollment_info + .and_then(|enrollment| enrollment.count), + start_date: protocol + .status_module + .start_date_struct + .map(|date| date.date), + completion_date: protocol + .status_module + .completion_date_struct + .map(|date| date.date), + } +} + +fn validate_nct_id(nct_id: &str) -> UseResult<()> { + let digits = nct_id.strip_prefix("NCT").unwrap_or_default(); + if digits.len() != 8 || !digits.bytes().all(|byte| byte.is_ascii_digit()) { + return Err(UseError::new( + "use.science.identifier_invalid", + "A ClinicalTrials.gov identifier must use the form NCT followed by eight digits.", + )); + } + Ok(()) +} + +fn validate_status(status: &str) -> UseResult<()> { + if status.is_empty() + || !status + .bytes() + .all(|byte| byte.is_ascii_uppercase() || byte == b'_') + { + return Err(UseError::new( + "use.science.input_invalid", + "Clinical trial statuses must use uppercase API values such as RECRUITING.", + )); + } + Ok(()) +} + +fn validate_page_token(token: &str) -> UseResult<()> { + if token.is_empty() + || token.len() > 2_048 + || !token + .bytes() + .all(|byte| byte.is_ascii_alphanumeric() || matches!(byte, b'-' | b'_' | b'.' | b'~')) + { + return Err(UseError::new( + "use.science.input_invalid", + "ClinicalTrials.gov page token contains unsupported characters.", + )); + } + Ok(()) +} + +fn bounded_limit(limit: usize) -> UseResult { + if !(1..=100).contains(&limit) { + return Err(UseError::new( + "use.science.limit_invalid", + "Clinical trial result limit must be between 1 and 100.", + )); + } + Ok(limit) +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn flattens_clinical_trial_modules() { + let study: Study = serde_json::from_value(serde_json::json!({ + "protocolSection": { + "identificationModule": { + "nctId": "NCT12345678", + "briefTitle": "Trial" + }, + "statusModule": { + "overallStatus": "RECRUITING", + "startDateStruct": {"date": "2026-01"} + }, + "designModule": { + "studyType": "INTERVENTIONAL", + "phases": ["PHASE2"], + "enrollmentInfo": {"count": 120} + }, + "conditionsModule": {"conditions": ["Cancer"]}, + "armsInterventionsModule": { + "interventions": [{"name": "Drug A"}] + }, + "sponsorCollaboratorsModule": { + "leadSponsor": {"name": "A3S Lab"} + } + } + })) + .unwrap(); + let trial = flatten_study(study); + assert_eq!(trial.nct_id, "NCT12345678"); + assert_eq!(trial.enrollment, Some(120)); + assert_eq!(trial.interventions, ["Drug A"]); + } + + #[test] + fn validates_trial_identifiers_and_statuses() { + assert_eq!( + validate_nct_id("NCT123").unwrap_err().code, + "use.science.identifier_invalid" + ); + assert_eq!( + validate_status("recruiting").unwrap_err().code, + "use.science.input_invalid" + ); + } +} diff --git a/crates/science/src/ensembl.rs b/crates/science/src/ensembl.rs new file mode 100644 index 00000000..586e4365 --- /dev/null +++ b/crates/science/src/ensembl.rs @@ -0,0 +1,203 @@ +use std::time::Duration; + +use a3s_use_core::{UseError, UseResult}; +use serde::Deserialize; + +use crate::models::{EnsemblGene, EnsemblHomolog}; +use crate::ScienceClient; + +const ENSEMBL_INTERVAL: Duration = Duration::from_millis(120); + +#[derive(Debug, Deserialize)] +struct GeneResponse { + id: String, + #[serde(default)] + display_name: Option, + #[serde(default)] + description: Option, + #[serde(default)] + species: Option, + #[serde(default)] + biotype: Option, + #[serde(default)] + seq_region_name: Option, + #[serde(default)] + start: Option, + #[serde(default)] + end: Option, + #[serde(default)] + strand: Option, +} + +#[derive(Debug, Deserialize)] +struct HomologyEnvelope { + #[serde(default)] + data: Vec, +} + +#[derive(Debug, Deserialize)] +struct HomologyData { + #[serde(default)] + id: Option, + #[serde(default)] + homologies: Vec, +} + +#[derive(Debug, Deserialize)] +struct RawHomology { + #[serde(default, rename = "type")] + homology_type: Option, + target: HomologyTarget, +} + +#[derive(Debug, Deserialize)] +struct HomologyTarget { + id: String, + #[serde(default)] + species: Option, + #[serde(default)] + protein_id: Option, + #[serde(default)] + perc_id: Option, + #[serde(default)] + perc_pos: Option, +} + +impl ScienceClient { + pub async fn ensembl_lookup_gene(&self, species: &str, symbol: &str) -> UseResult { + validate_species(species)?; + validate_symbol(symbol)?; + let url = self.endpoint_url( + &self.endpoints.ensembl, + &["lookup", "symbol", species, symbol], + )?; + let response: GeneResponse = self + .get_json( + "Ensembl", + self.http + .get(url) + .header(reqwest::header::ACCEPT, "application/json"), + ENSEMBL_INTERVAL, + ) + .await?; + Ok(EnsemblGene { + id: response.id, + display_name: response.display_name, + description: response.description, + species: response.species, + biotype: response.biotype, + chromosome: response.seq_region_name, + start: response.start, + end: response.end, + strand: response.strand, + }) + } + + pub async fn ensembl_homologs( + &self, + species: &str, + symbol: &str, + target_species: Option<&str>, + limit: usize, + ) -> UseResult> { + validate_species(species)?; + validate_symbol(symbol)?; + if let Some(target_species) = target_species { + validate_species(target_species)?; + } + if !(1..=200).contains(&limit) { + return Err(UseError::new( + "use.science.limit_invalid", + "Ensembl homolog result limit must be between 1 and 200.", + )); + } + let url = self.endpoint_url( + &self.endpoints.ensembl, + &["homology", "symbol", species, symbol], + )?; + let mut query = vec![ + ("type", "orthologues".to_string()), + ("format", "condensed".to_string()), + ]; + if let Some(target_species) = target_species { + query.push(("target_species", target_species.to_string())); + } + let envelope: HomologyEnvelope = self + .get_json( + "Ensembl", + self.http + .get(url) + .query(&query) + .header(reqwest::header::ACCEPT, "application/json"), + ENSEMBL_INTERVAL, + ) + .await?; + let mut items = Vec::new(); + for data in envelope.data { + for homology in data.homologies { + items.push(EnsemblHomolog { + homology_type: homology.homology_type, + source_gene_id: data.id.clone(), + target_gene_id: homology.target.id, + target_species: homology.target.species, + target_protein_id: homology.target.protein_id, + identity_percent: homology.target.perc_id, + positive_percent: homology.target.perc_pos, + }); + if items.len() == limit { + return Ok(items); + } + } + } + Ok(items) + } +} + +fn validate_species(species: &str) -> UseResult<()> { + if species.is_empty() + || species.len() > 100 + || !species + .bytes() + .all(|byte| byte.is_ascii_lowercase() || byte == b'_') + { + return Err(UseError::new( + "use.science.input_invalid", + "Ensembl species must use a lowercase identifier such as homo_sapiens.", + )); + } + Ok(()) +} + +fn validate_symbol(symbol: &str) -> UseResult<()> { + if symbol.is_empty() + || symbol.len() > 100 + || !symbol + .bytes() + .all(|byte| byte.is_ascii_alphanumeric() || matches!(byte, b'.' | b'-' | b'_')) + { + return Err(UseError::new( + "use.science.input_invalid", + "Ensembl gene symbol contains unsupported characters.", + )); + } + Ok(()) +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn validates_species_and_symbols_without_accepting_paths() { + assert!(validate_species("homo_sapiens").is_ok()); + assert!(validate_symbol("TP53").is_ok()); + assert_eq!( + validate_species("../human").unwrap_err().code, + "use.science.input_invalid" + ); + assert_eq!( + validate_symbol("TP53/../../x").unwrap_err().code, + "use.science.input_invalid" + ); + } +} diff --git a/crates/science/src/lib.rs b/crates/science/src/lib.rs new file mode 100644 index 00000000..92e2b5f7 --- /dev/null +++ b/crates/science/src/lib.rs @@ -0,0 +1,24 @@ +//! Typed, read-only life-science data retrieval for A3S Use. +//! +//! The crate deliberately exposes public upstream contracts rather than a +//! generic action envelope. [ScienceClient] is suitable for embedding, while +//! [ScienceMcpServer] presents the same operations as standard MCP tools. + +mod biorxiv; +mod chembl; +pub mod cli; +mod client; +mod clinical_trials; +mod ensembl; +pub mod mcp; +mod models; +mod pubmed; + +pub use client::{ScienceClient, ScienceClientBuilder, ScienceEndpoints}; +pub use mcp::ScienceMcpServer; +pub use models::{ + BioRxivPage, BioRxivRecord, ChemblActivity, ChemblMolecule, ChemblTarget, ClinicalTrial, + ClinicalTrialPage, EnsemblGene, EnsemblHomolog, Page, PubMedArticle, ScienceDiagnostic, +}; + +pub use a3s_use_core::{UseError, UseResult}; diff --git a/crates/science/src/main.rs b/crates/science/src/main.rs new file mode 100644 index 00000000..7f2e37f8 --- /dev/null +++ b/crates/science/src/main.rs @@ -0,0 +1,39 @@ +use std::process::ExitCode; + +#[tokio::main] +async fn main() -> ExitCode { + let args = std::env::args().skip(1).collect::>(); + let json = args.iter().any(|argument| argument == "--json"); + match a3s_use_science::cli::run(args).await { + Ok(output) => { + if output.should_print && json { + println!( + "{}", + serde_json::to_string_pretty(&output.json).unwrap_or_default() + ); + } else if output.should_print && !output.human.is_empty() { + println!("{}", output.human); + } + ExitCode::from(output.exit_code) + } + Err(error) => { + if json { + let output = serde_json::json!({ + "schemaVersion": 1, + "ok": false, + "error": error, + }); + println!( + "{}", + serde_json::to_string_pretty(&output).unwrap_or_default() + ); + } else { + eprintln!("a3s-use-science: {error}"); + if let Some(suggestion) = &error.suggestion { + eprintln!("suggestion: {suggestion}"); + } + } + ExitCode::from(1) + } + } +} diff --git a/crates/science/src/mcp.rs b/crates/science/src/mcp.rs new file mode 100644 index 00000000..123403b8 --- /dev/null +++ b/crates/science/src/mcp.rs @@ -0,0 +1,434 @@ +//! Standard MCP tools for the process-isolated Science extension. + +use rmcp::handler::server::{router::tool::ToolRouter, wrapper::Parameters}; +use rmcp::model::{CallToolResult, Implementation, ServerCapabilities, ServerInfo}; +use rmcp::{tool, tool_handler, tool_router, ServerHandler, ServiceExt}; +use serde::{Deserialize, Serialize}; + +use crate::{ScienceClient, UseError, UseResult}; + +const DEFAULT_LIMIT: usize = 20; + +#[derive(Debug, Clone, Deserialize, schemars::JsonSchema)] +#[serde(rename_all = "camelCase")] +struct PubMedSearchInput { + #[schemars(description = "PubMed search expression")] + query: String, + #[schemars(description = "Maximum article summaries to return; defaults to 20, maximum 100")] + limit: Option, +} + +#[derive(Debug, Clone, Deserialize, schemars::JsonSchema)] +struct PubMedGetInput { + #[schemars(description = "PubMed identifier containing 1 to 12 digits")] + pmid: String, +} + +#[derive(Debug, Clone, Deserialize, schemars::JsonSchema)] +#[serde(rename_all = "camelCase")] +struct ChemblSearchInput { + #[schemars(description = "Free-text ChEMBL search expression")] + query: String, + #[schemars(description = "Maximum records to return; defaults to 20, maximum 100")] + limit: Option, +} + +#[derive(Debug, Clone, Deserialize, schemars::JsonSchema)] +#[serde(rename_all = "camelCase")] +struct ChemblMoleculeInput { + #[schemars(description = "ChEMBL molecule identifier, such as CHEMBL25")] + chembl_id: String, +} + +#[derive(Debug, Clone, Deserialize, schemars::JsonSchema)] +#[serde(rename_all = "camelCase")] +struct ChemblActivitiesInput { + #[schemars(description = "Optional ChEMBL molecule identifier")] + molecule_chembl_id: Option, + #[schemars(description = "Optional ChEMBL target identifier")] + target_chembl_id: Option, + #[schemars(description = "Maximum activities to return; defaults to 20, maximum 100")] + limit: Option, +} + +#[derive(Debug, Clone, Deserialize, schemars::JsonSchema)] +#[serde(rename_all = "camelCase")] +struct ClinicalTrialsSearchInput { + #[schemars(description = "ClinicalTrials.gov query expression")] + query: String, + #[schemars(description = "Optional uppercase statuses, such as RECRUITING")] + statuses: Option>, + #[schemars(description = "Maximum studies to return; defaults to 20, maximum 100")] + limit: Option, + #[schemars(description = "Opaque next-page token from a prior response")] + page_token: Option, +} + +#[derive(Debug, Clone, Deserialize, schemars::JsonSchema)] +#[serde(rename_all = "camelCase")] +struct ClinicalTrialGetInput { + #[schemars(description = "ClinicalTrials.gov identifier, such as NCT01234567")] + nct_id: String, +} + +#[derive(Debug, Clone, Deserialize, schemars::JsonSchema)] +#[serde(rename_all = "camelCase")] +struct BioRxivSearchInput { + #[schemars(description = "Inclusive range start in YYYY-MM-DD form")] + from_date: String, + #[schemars(description = "Inclusive range end in YYYY-MM-DD form")] + to_date: String, + #[schemars(description = "Optional case-insensitive title, author, abstract, or DOI filter")] + query: Option, + #[schemars(description = "Optional exact bioRxiv category filter")] + category: Option, + #[schemars(description = "Maximum preprints to return; defaults to 20, maximum 100")] + limit: Option, +} + +#[derive(Debug, Clone, Deserialize, schemars::JsonSchema)] +struct BioRxivGetInput { + #[schemars(description = "bioRxiv DOI beginning with 10.1101/")] + doi: String, +} + +#[derive(Debug, Clone, Deserialize, schemars::JsonSchema)] +struct EnsemblLookupInput { + #[schemars(description = "Lowercase Ensembl species identifier, such as homo_sapiens")] + species: String, + #[schemars(description = "Gene symbol, such as TP53")] + symbol: String, +} + +#[derive(Debug, Clone, Deserialize, schemars::JsonSchema)] +#[serde(rename_all = "camelCase")] +struct EnsemblHomologsInput { + #[schemars(description = "Lowercase source species identifier")] + species: String, + #[schemars(description = "Source gene symbol")] + symbol: String, + #[schemars(description = "Optional lowercase target species identifier")] + target_species: Option, + #[schemars(description = "Maximum homologs to return; defaults to 50, maximum 200")] + limit: Option, +} + +#[derive(Clone)] +pub struct ScienceMcpServer { + client: ScienceClient, + tool_router: ToolRouter, +} + +impl ScienceMcpServer { + pub fn new(client: ScienceClient) -> Self { + Self { + client, + tool_router: Self::tool_router(), + } + } + + pub fn from_env() -> UseResult { + Ok(Self::new(ScienceClient::from_env()?)) + } + + /// Serve standard MCP framing over stdin/stdout until the peer disconnects. + pub async fn serve_stdio(self) -> UseResult<()> { + let service = self + .serve(rmcp::transport::stdio()) + .await + .map_err(|error| mcp_error("start", error))?; + service + .waiting() + .await + .map_err(|error| mcp_error("run", error))?; + Ok(()) + } +} + +#[tool_router] +impl ScienceMcpServer { + #[tool( + name = "science_doctor", + description = "Inspect Science extension configuration without making a network request" + )] + async fn science_doctor(&self) -> Result { + Ok(tool_result(Ok(self.client.diagnostic()))) + } + + #[tool( + name = "science_pubmed_search", + description = "Search PubMed and return typed article summaries with stable identifiers" + )] + async fn science_pubmed_search( + &self, + Parameters(input): Parameters, + ) -> Result { + Ok(tool_result( + self.client + .pubmed_search(&input.query, input.limit.unwrap_or(DEFAULT_LIMIT)) + .await, + )) + } + + #[tool( + name = "science_pubmed_get", + description = "Retrieve one PubMed article summary by PMID" + )] + async fn science_pubmed_get( + &self, + Parameters(input): Parameters, + ) -> Result { + Ok(tool_result(self.client.pubmed_get(&input.pmid).await)) + } + + #[tool( + name = "science_chembl_search_molecules", + description = "Search ChEMBL molecules and return normalized identifiers and structures" + )] + async fn science_chembl_search_molecules( + &self, + Parameters(input): Parameters, + ) -> Result { + Ok(tool_result( + self.client + .chembl_search_molecules(&input.query, input.limit.unwrap_or(DEFAULT_LIMIT)) + .await, + )) + } + + #[tool( + name = "science_chembl_get_molecule", + description = "Retrieve one normalized ChEMBL molecule record" + )] + async fn science_chembl_get_molecule( + &self, + Parameters(input): Parameters, + ) -> Result { + Ok(tool_result( + self.client.chembl_get_molecule(&input.chembl_id).await, + )) + } + + #[tool( + name = "science_chembl_search_targets", + description = "Search ChEMBL targets and return normalized target records" + )] + async fn science_chembl_search_targets( + &self, + Parameters(input): Parameters, + ) -> Result { + Ok(tool_result( + self.client + .chembl_search_targets(&input.query, input.limit.unwrap_or(DEFAULT_LIMIT)) + .await, + )) + } + + #[tool( + name = "science_chembl_activities", + description = "Retrieve ChEMBL bioactivity records for a molecule, target, or both" + )] + async fn science_chembl_activities( + &self, + Parameters(input): Parameters, + ) -> Result { + Ok(tool_result( + self.client + .chembl_activities( + input.molecule_chembl_id.as_deref(), + input.target_chembl_id.as_deref(), + input.limit.unwrap_or(DEFAULT_LIMIT), + ) + .await, + )) + } + + #[tool( + name = "science_clinical_trials_search", + description = "Search ClinicalTrials.gov and return normalized protocol summaries" + )] + async fn science_clinical_trials_search( + &self, + Parameters(input): Parameters, + ) -> Result { + Ok(tool_result( + self.client + .clinical_trials_search( + &input.query, + input.statuses.as_deref().unwrap_or_default(), + input.limit.unwrap_or(DEFAULT_LIMIT), + input.page_token.as_deref(), + ) + .await, + )) + } + + #[tool( + name = "science_clinical_trial_get", + description = "Retrieve one ClinicalTrials.gov study by NCT identifier" + )] + async fn science_clinical_trial_get( + &self, + Parameters(input): Parameters, + ) -> Result { + Ok(tool_result( + self.client.clinical_trial_get(&input.nct_id).await, + )) + } + + #[tool( + name = "science_biorxiv_search", + description = "Search a bounded bioRxiv date range and return matching preprints" + )] + async fn science_biorxiv_search( + &self, + Parameters(input): Parameters, + ) -> Result { + Ok(tool_result( + self.client + .biorxiv_search( + &input.from_date, + &input.to_date, + input.query.as_deref(), + input.category.as_deref(), + input.limit.unwrap_or(DEFAULT_LIMIT), + ) + .await, + )) + } + + #[tool( + name = "science_biorxiv_get", + description = "Retrieve all returned versions of one bioRxiv DOI" + )] + async fn science_biorxiv_get( + &self, + Parameters(input): Parameters, + ) -> Result { + Ok(tool_result(self.client.biorxiv_get(&input.doi).await)) + } + + #[tool( + name = "science_ensembl_lookup_gene", + description = "Look up an Ensembl gene by species and symbol" + )] + async fn science_ensembl_lookup_gene( + &self, + Parameters(input): Parameters, + ) -> Result { + Ok(tool_result( + self.client + .ensembl_lookup_gene(&input.species, &input.symbol) + .await, + )) + } + + #[tool( + name = "science_ensembl_homologs", + description = "Retrieve Ensembl orthologs for one gene symbol" + )] + async fn science_ensembl_homologs( + &self, + Parameters(input): Parameters, + ) -> Result { + Ok(tool_result( + self.client + .ensembl_homologs( + &input.species, + &input.symbol, + input.target_species.as_deref(), + input.limit.unwrap_or(50), + ) + .await, + )) + } +} + +#[tool_handler] +impl ServerHandler for ScienceMcpServer { + fn get_info(&self) -> ServerInfo { + ServerInfo { + capabilities: ServerCapabilities::builder().enable_tools().build(), + server_info: Implementation { + name: "a3s-use-science".to_string(), + title: Some("A3S Use Science".to_string()), + version: env!("CARGO_PKG_VERSION").to_string(), + icons: None, + website_url: Some("https://github.com/A3S-Lab/Use".to_string()), + }, + instructions: Some( + "Use science_doctor before retrieval. Preserve returned source identifiers, distinguish bioRxiv preprints from peer-reviewed literature, and do not present public database records as medical advice. PubMed requires A3S_SCIENCE_CONTACT_EMAIL." + .to_string(), + ), + ..Default::default() + } + } +} + +fn tool_result(result: UseResult) -> CallToolResult +where + T: Serialize, +{ + match result { + Ok(output) => match serde_json::to_value(output) { + Ok(value) => CallToolResult::structured(value), + Err(error) => tool_error(UseError::new( + "use.science.output_invalid", + format!("Failed to encode Science MCP output: {error}"), + )), + }, + Err(error) => tool_error(error), + } +} + +fn tool_error(error: UseError) -> CallToolResult { + CallToolResult::structured_error(serde_json::to_value(error).unwrap_or_else(|_| { + serde_json::json!({ + "code": "use.error_encoding_failed", + "message": "Failed to encode A3S Use error." + }) + })) +} + +fn mcp_error(action: &str, error: impl std::fmt::Display) -> UseError { + UseError::new( + "use.science.mcp_failed", + format!("Failed to {action} the Science MCP server: {error}"), + ) +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn server_exposes_only_typed_read_tools() { + let client = ScienceClient::builder().build().unwrap(); + let server = ScienceMcpServer::new(client); + let mut names = server + .tool_router + .list_all() + .iter() + .map(|tool| tool.name.to_string()) + .collect::>(); + names.sort_unstable(); + assert_eq!( + names, + [ + "science_biorxiv_get", + "science_biorxiv_search", + "science_chembl_activities", + "science_chembl_get_molecule", + "science_chembl_search_molecules", + "science_chembl_search_targets", + "science_clinical_trial_get", + "science_clinical_trials_search", + "science_doctor", + "science_ensembl_homologs", + "science_ensembl_lookup_gene", + "science_pubmed_get", + "science_pubmed_search", + ] + ); + } +} diff --git a/crates/science/src/models.rs b/crates/science/src/models.rs new file mode 100644 index 00000000..dcf48cc5 --- /dev/null +++ b/crates/science/src/models.rs @@ -0,0 +1,130 @@ +use serde::{Deserialize, Serialize}; + +#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] +#[serde(rename_all = "camelCase")] +pub struct Page { + pub total: Option, + pub next_page_token: Option, + pub items: Vec, +} + +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +#[serde(rename_all = "camelCase")] +pub struct PubMedArticle { + pub pmid: String, + pub title: String, + pub authors: Vec, + pub journal: Option, + pub publication_date: Option, + pub doi: Option, +} + +#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] +#[serde(rename_all = "camelCase")] +pub struct ChemblMolecule { + pub chembl_id: String, + pub preferred_name: Option, + pub molecule_type: Option, + pub max_phase: Option, + pub canonical_smiles: Option, + pub standard_inchi_key: Option, +} + +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +#[serde(rename_all = "camelCase")] +pub struct ChemblTarget { + pub chembl_id: String, + pub preferred_name: Option, + pub target_type: Option, + pub organism: Option, +} + +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +#[serde(rename_all = "camelCase")] +pub struct ChemblActivity { + pub activity_id: Option, + pub molecule_chembl_id: Option, + pub target_chembl_id: Option, + pub assay_chembl_id: Option, + pub standard_type: Option, + pub standard_relation: Option, + pub standard_value: Option, + pub standard_units: Option, + pub pchembl_value: Option, +} + +#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] +#[serde(rename_all = "camelCase")] +pub struct ClinicalTrial { + pub nct_id: String, + pub brief_title: String, + pub official_title: Option, + pub overall_status: Option, + pub study_type: Option, + pub phases: Vec, + pub conditions: Vec, + pub interventions: Vec, + pub lead_sponsor: Option, + pub enrollment: Option, + pub start_date: Option, + pub completion_date: Option, +} + +pub type ClinicalTrialPage = Page; + +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +#[serde(rename_all = "camelCase")] +pub struct BioRxivRecord { + pub doi: String, + pub title: String, + pub authors: String, + pub abstract_text: Option, + pub category: Option, + pub date: Option, + pub version: Option, + pub published_doi: Option, +} + +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +#[serde(rename_all = "camelCase")] +pub struct BioRxivPage { + pub total_upstream: Option, + pub scanned: usize, + pub items: Vec, +} + +#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] +#[serde(rename_all = "camelCase")] +pub struct EnsemblGene { + pub id: String, + pub display_name: Option, + pub description: Option, + pub species: Option, + pub biotype: Option, + pub chromosome: Option, + pub start: Option, + pub end: Option, + pub strand: Option, +} + +#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] +#[serde(rename_all = "camelCase")] +pub struct EnsemblHomolog { + pub homology_type: Option, + pub source_gene_id: Option, + pub target_gene_id: String, + pub target_species: Option, + pub target_protein_id: Option, + pub identity_percent: Option, + pub positive_percent: Option, +} + +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +#[serde(rename_all = "camelCase")] +pub struct ScienceDiagnostic { + pub version: String, + pub contact_email_configured: bool, + pub ncbi_api_key_configured: bool, + pub sources: Vec, + pub message: String, +} diff --git a/crates/science/src/pubmed.rs b/crates/science/src/pubmed.rs new file mode 100644 index 00000000..d31edd30 --- /dev/null +++ b/crates/science/src/pubmed.rs @@ -0,0 +1,248 @@ +use std::time::Duration; + +use a3s_use_core::{UseError, UseResult}; +use serde::Deserialize; +use serde_json::Value; + +use crate::models::{Page, PubMedArticle}; +use crate::ScienceClient; + +const PUBMED_KEYLESS_INTERVAL: Duration = Duration::from_millis(350); +const PUBMED_KEYED_INTERVAL: Duration = Duration::from_millis(110); + +#[derive(Debug, Deserialize)] +struct SearchEnvelope { + esearchresult: SearchResult, +} + +#[derive(Debug, Deserialize)] +struct SearchResult { + count: String, + #[serde(default)] + idlist: Vec, +} + +#[derive(Debug, Deserialize)] +struct SummaryEnvelope { + result: serde_json::Map, +} + +impl ScienceClient { + pub async fn pubmed_search(&self, query: &str, limit: usize) -> UseResult> { + let query = required_text("PubMed query", query)?; + let limit = bounded_limit(limit, 100)?; + let contact_email = self.pubmed_contact_email()?; + let url = self.endpoint_url(&self.endpoints.pubmed, &["esearch.fcgi"])?; + let mut params = vec![ + ("db", "pubmed".to_string()), + ("term", query.to_string()), + ("retmode", "json".to_string()), + ("retmax", limit.to_string()), + ("tool", "a3s-use-science".to_string()), + ("email", contact_email.to_string()), + ]; + if let Some(api_key) = &self.ncbi_api_key { + params.push(("api_key", api_key.clone())); + } + let envelope: SearchEnvelope = self + .get_json( + "PubMed", + self.http.get(url).query(¶ms), + self.pubmed_interval(), + ) + .await?; + let total = envelope.esearchresult.count.parse().ok(); + if envelope.esearchresult.idlist.is_empty() { + return Ok(Page { + total, + next_page_token: None, + items: Vec::new(), + }); + } + let items = self + .pubmed_summaries(&envelope.esearchresult.idlist) + .await?; + Ok(Page { + total, + next_page_token: None, + items, + }) + } + + pub async fn pubmed_get(&self, pmid: &str) -> UseResult { + validate_pmid(pmid)?; + self.pubmed_summaries(&[pmid.to_string()]) + .await? + .into_iter() + .next() + .ok_or_else(|| { + UseError::new( + "use.science.not_found", + format!("PubMed did not return PMID {pmid}."), + ) + .with_detail("service", "PubMed") + .with_detail("pmid", pmid) + }) + } + + async fn pubmed_summaries(&self, pmids: &[String]) -> UseResult> { + let contact_email = self.pubmed_contact_email()?; + let url = self.endpoint_url(&self.endpoints.pubmed, &["esummary.fcgi"])?; + let mut params = vec![ + ("db", "pubmed".to_string()), + ("id", pmids.join(",")), + ("retmode", "json".to_string()), + ("version", "2.0".to_string()), + ("tool", "a3s-use-science".to_string()), + ("email", contact_email.to_string()), + ]; + if let Some(api_key) = &self.ncbi_api_key { + params.push(("api_key", api_key.clone())); + } + let envelope: SummaryEnvelope = self + .get_json( + "PubMed", + self.http.get(url).query(¶ms), + self.pubmed_interval(), + ) + .await?; + Ok(parse_summaries(pmids, &envelope.result)) + } + + fn pubmed_contact_email(&self) -> UseResult<&str> { + self.contact_email.as_deref().ok_or_else(|| { + UseError::new( + "use.science.contact_email_required", + "PubMed requests require a contact email for responsible NCBI E-utilities use.", + ) + .with_suggestion("Set A3S_SCIENCE_CONTACT_EMAIL and retry.") + }) + } + + fn pubmed_interval(&self) -> Duration { + if self.ncbi_api_key.is_some() { + PUBMED_KEYED_INTERVAL + } else { + PUBMED_KEYLESS_INTERVAL + } + } +} + +fn parse_summaries( + pmids: &[String], + result: &serde_json::Map, +) -> Vec { + pmids + .iter() + .filter_map(|pmid| { + let record = result.get(pmid)?.as_object()?; + let authors = record + .get("authors") + .and_then(Value::as_array) + .into_iter() + .flatten() + .filter_map(|author| author.get("name").and_then(Value::as_str)) + .map(str::to_string) + .collect(); + let doi = record + .get("articleids") + .and_then(Value::as_array) + .into_iter() + .flatten() + .find(|identifier| identifier.get("idtype").and_then(Value::as_str) == Some("doi")) + .and_then(|identifier| identifier.get("value")) + .and_then(Value::as_str) + .map(str::to_string); + Some(PubMedArticle { + pmid: pmid.clone(), + title: string_field(record, "title").unwrap_or_default(), + authors, + journal: string_field(record, "fulljournalname") + .or_else(|| string_field(record, "source")), + publication_date: string_field(record, "pubdate"), + doi, + }) + }) + .collect() +} + +fn string_field(record: &serde_json::Map, name: &str) -> Option { + record + .get(name) + .and_then(Value::as_str) + .filter(|value| !value.is_empty()) + .map(str::to_string) +} + +fn required_text<'a>(label: &str, value: &'a str) -> UseResult<&'a str> { + let value = value.trim(); + if value.is_empty() { + return Err(UseError::new( + "use.science.input_invalid", + format!("{label} cannot be empty."), + )); + } + Ok(value) +} + +fn bounded_limit(limit: usize, maximum: usize) -> UseResult { + if !(1..=maximum).contains(&limit) { + return Err(UseError::new( + "use.science.limit_invalid", + format!("Result limit must be between 1 and {maximum}."), + )); + } + Ok(limit) +} + +fn validate_pmid(pmid: &str) -> UseResult<()> { + if pmid.is_empty() || pmid.len() > 12 || !pmid.bytes().all(|byte| byte.is_ascii_digit()) { + return Err(UseError::new( + "use.science.identifier_invalid", + "A PMID must contain only 1 to 12 digits.", + )); + } + Ok(()) +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn parses_pubmed_summaries_in_requested_order() { + let result = serde_json::json!({ + "2": { + "title": "Second", + "authors": [{"name": "B Author"}], + "articleids": [{"idtype": "doi", "value": "10.1/second"}] + }, + "1": { + "title": "First", + "authors": [{"name": "A Author"}], + "fulljournalname": "Journal", + "pubdate": "2026", + "articleids": [] + } + }); + let articles = parse_summaries( + &["1".to_string(), "2".to_string()], + result.as_object().unwrap(), + ); + assert_eq!(articles[0].pmid, "1"); + assert_eq!(articles[0].authors, ["A Author"]); + assert_eq!(articles[1].doi.as_deref(), Some("10.1/second")); + } + + #[test] + fn validates_pubmed_inputs() { + assert_eq!( + validate_pmid("../1").unwrap_err().code, + "use.science.identifier_invalid" + ); + assert_eq!( + bounded_limit(0, 100).unwrap_err().code, + "use.science.limit_invalid" + ); + } +} diff --git a/crates/science/tests/integration.rs b/crates/science/tests/integration.rs new file mode 100644 index 00000000..9e203a98 --- /dev/null +++ b/crates/science/tests/integration.rs @@ -0,0 +1,318 @@ +use std::collections::HashMap; +use std::path::Path; +use std::process::Command; +use std::sync::{Arc, Mutex}; + +use a3s_use_core::RiskClass; +use a3s_use_extension::{ + ExtensionManifest, ExtensionPaths, ExtensionRegistry, InstallOptions, McpTransport, +}; +use a3s_use_science::{ScienceClient, ScienceEndpoints}; +use axum::extract::{OriginalUri, Query, State}; +use axum::http::{HeaderMap, StatusCode, Uri}; +use axum::response::IntoResponse; +use axum::routing::get; +use axum::{Json, Router}; +use serde_json::json; +use tokio::task::JoinHandle; +use url::Url; + +#[derive(Clone, Default)] +struct RequestLog(Arc>>); + +#[derive(Debug)] +struct RecordedRequest { + uri: Uri, + query: HashMap, + user_agent: Option, +} + +struct MockServer { + base: Url, + log: RequestLog, + task: JoinHandle<()>, +} + +impl MockServer { + async fn start() -> Self { + let log = RequestLog::default(); + let app = Router::new() + .route("/pubmed/esearch.fcgi", get(pubmed_search)) + .route("/pubmed/esummary.fcgi", get(pubmed_summary)) + .route("/chembl/molecule/search.json", get(chembl_failure)) + .with_state(log.clone()); + let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap(); + let address = listener.local_addr().unwrap(); + let task = tokio::spawn(async move { + let _ = axum::serve(listener, app).await; + }); + Self { + base: Url::parse(&format!("http://{address}/")).unwrap(), + log, + task, + } + } + + fn endpoints(&self) -> ScienceEndpoints { + ScienceEndpoints { + pubmed: self.base.join("pubmed/").unwrap(), + chembl: self.base.join("chembl/").unwrap(), + clinical_trials: self.base.join("clinical-trials/").unwrap(), + biorxiv: self.base.join("biorxiv/").unwrap(), + ensembl: self.base.join("ensembl/").unwrap(), + } + } +} + +impl Drop for MockServer { + fn drop(&mut self) { + self.task.abort(); + } +} + +async fn pubmed_search( + State(log): State, + OriginalUri(uri): OriginalUri, + headers: HeaderMap, + Query(query): Query>, +) -> Json { + record(&log, uri, query, &headers); + Json(json!({ + "esearchresult": { + "count": "1", + "idlist": ["12345678"] + } + })) +} + +async fn pubmed_summary( + State(log): State, + OriginalUri(uri): OriginalUri, + headers: HeaderMap, + Query(query): Query>, +) -> Json { + record(&log, uri, query, &headers); + Json(json!({ + "result": { + "uids": ["12345678"], + "12345678": { + "title": "A typed science result", + "authors": [{"name": "A. Researcher"}], + "fulljournalname": "Journal of Tests", + "pubdate": "2026", + "articleids": [{"idtype": "doi", "value": "10.1000/test"}] + } + } + })) +} + +async fn chembl_failure( + State(log): State, + OriginalUri(uri): OriginalUri, + headers: HeaderMap, + Query(query): Query>, +) -> impl IntoResponse { + record(&log, uri, query, &headers); + ( + StatusCode::SERVICE_UNAVAILABLE, + "upstream failure ".repeat(100), + ) +} + +fn record(log: &RequestLog, uri: Uri, query: HashMap, headers: &HeaderMap) { + let user_agent = headers + .get(axum::http::header::USER_AGENT) + .and_then(|value| value.to_str().ok()) + .map(str::to_string); + log.0.lock().unwrap().push(RecordedRequest { + uri, + query, + user_agent, + }); +} + +#[tokio::test] +async fn pubmed_uses_two_typed_requests_and_encodes_contact_metadata() { + let server = MockServer::start().await; + let client = ScienceClient::builder() + .endpoints(server.endpoints()) + .contact_email("researcher@example.org") + .ncbi_api_key("test-key") + .build() + .unwrap(); + + let page = client + .pubmed_search("gene therapy & safety", 7) + .await + .unwrap(); + assert_eq!(page.total, Some(1)); + assert_eq!(page.items[0].pmid, "12345678"); + assert_eq!(page.items[0].doi.as_deref(), Some("10.1000/test")); + + let requests = server.log.0.lock().unwrap(); + assert_eq!(requests.len(), 2); + assert!(requests[0].uri.to_string().contains("%26")); + assert_eq!(requests[0].query["term"], "gene therapy & safety"); + assert_eq!(requests[0].query["retmax"], "7"); + assert_eq!(requests[0].query["email"], "researcher@example.org"); + assert_eq!(requests[0].query["api_key"], "test-key"); + assert_eq!( + requests[0].user_agent.as_deref(), + Some(concat!("a3s-use-science/", env!("CARGO_PKG_VERSION"))) + ); + assert_eq!(requests[1].query["id"], "12345678"); +} + +#[tokio::test] +async fn upstream_http_failures_use_a_stable_bounded_error() { + let server = MockServer::start().await; + let client = ScienceClient::builder() + .endpoints(server.endpoints()) + .build() + .unwrap(); + + let error = client + .chembl_search_molecules("aspirin", 3) + .await + .unwrap_err(); + assert_eq!(error.code, "use.science.upstream_error"); + assert_eq!(error.details["service"], "ChEMBL"); + assert_eq!(error.details["status"], 503); + let body = error.details["body"].as_str().unwrap(); + assert_eq!(body.chars().count(), 1_025); + assert!(body.ends_with('…')); + + let requests = server.log.0.lock().unwrap(); + assert_eq!(requests[0].query["q"], "aspirin"); + assert_eq!(requests[0].query["limit"], "3"); +} + +#[test] +fn packaged_manifest_declares_native_read_only_surfaces() { + let manifest_text = include_str!("../package/a3s-use-extension.acl"); + let manifest = ExtensionManifest::parse_acl(manifest_text).unwrap(); + assert_eq!(manifest.package_id, "a3s/science"); + assert_eq!(manifest.version, env!("CARGO_PKG_VERSION")); + assert_eq!(manifest.route, "science"); + assert_eq!(manifest.actions, [RiskClass::Read]); + assert!(manifest.cli.as_ref().unwrap().json_output); + assert_eq!( + manifest.mcp.as_ref().unwrap().transport, + McpTransport::Stdio + ); + assert_eq!( + manifest.mcp.as_ref().unwrap().args, + ["serve".to_string(), "--mcp".to_string()] + ); + assert_eq!( + manifest.skill.as_ref().unwrap().path, + Path::new("skills/a3s-use-science/SKILL.md") + ); + manifest + .validate_package_root( + Path::new(env!("CARGO_MANIFEST_DIR")) + .join("package") + .as_path(), + ) + .unwrap(); +} + +#[test] +fn binary_emits_versioned_diagnostics_and_errors() { + let binary = env!("CARGO_BIN_EXE_a3s-use-science"); + let diagnostic = Command::new(binary) + .args(["doctor", "--json"]) + .output() + .unwrap(); + assert!(diagnostic.status.success()); + let value: serde_json::Value = serde_json::from_slice(&diagnostic.stdout).unwrap(); + assert_eq!(value["schemaVersion"], 1); + assert_eq!(value["data"]["sources"].as_array().unwrap().len(), 5); + + let invalid = Command::new(binary) + .args(["pubmed", "get", "../escape", "--json"]) + .output() + .unwrap(); + assert!(!invalid.status.success()); + let value: serde_json::Value = serde_json::from_slice(&invalid.stdout).unwrap(); + assert_eq!(value["error"]["code"], "use.science.identifier_invalid"); +} + +#[tokio::test] +async fn real_science_package_installs_hot_upgrades_dispatches_and_uninstalls() { + let temp = tempfile::tempdir().unwrap(); + let first_package = temp.path().join("science-package-v1"); + let second_package = temp.path().join("science-package-v2"); + create_science_package(&first_package); + create_science_package(&second_package); + let registry = ExtensionRegistry::new(ExtensionPaths::new( + temp.path().join("data"), + temp.path().join("state"), + )); + + let installed = registry + .install_local( + "a3s/science", + &first_package, + InstallOptions { + allow_unsigned: true, + ..InstallOptions::default() + }, + ) + .await + .unwrap(); + assert!(installed.changed); + let first_root = installed.extension.receipt.package_root.clone(); + let lease = registry.acquire_route("science").await.unwrap().unwrap(); + let executable = lease.extension().cli_executable().unwrap(); + let diagnostic = Command::new(executable) + .args(["doctor", "--json"]) + .output() + .unwrap(); + assert!(diagnostic.status.success()); + let value: serde_json::Value = serde_json::from_slice(&diagnostic.stdout).unwrap(); + assert_eq!(value["schemaVersion"], 1); + assert_eq!(value["data"]["sources"].as_array().unwrap().len(), 5); + + let upgraded = registry + .install_local( + "a3s/science", + &second_package, + InstallOptions { + force: true, + allow_unsigned: true, + }, + ) + .await + .unwrap(); + assert!(upgraded.changed); + assert_ne!(upgraded.extension.receipt.package_root, first_root); + assert!( + first_root.exists(), + "the generation pinned by an active route lease was removed" + ); + drop(lease); + + let removed = registry.uninstall("a3s/science").await.unwrap(); + assert!(removed.changed); + assert!(registry.get("a3s/science").await.unwrap().is_none()); + assert!(!first_root.parent().unwrap().exists()); +} + +fn create_science_package(root: &Path) { + let binary = root.join("bin/a3s-use-science"); + let skill = root.join("skills/a3s-use-science/SKILL.md"); + std::fs::create_dir_all(binary.parent().unwrap()).unwrap(); + std::fs::create_dir_all(skill.parent().unwrap()).unwrap(); + std::fs::copy(env!("CARGO_BIN_EXE_a3s-use-science"), &binary).unwrap(); + std::fs::write( + root.join("a3s-use-extension.acl"), + include_str!("../package/a3s-use-extension.acl"), + ) + .unwrap(); + std::fs::write( + &skill, + include_str!("../package/skills/a3s-use-science/SKILL.md"), + ) + .unwrap(); +} diff --git a/docs/architecture.md b/docs/architecture.md index c152127b..c3238948 100644 --- a/docs/architecture.md +++ b/docs/architecture.md @@ -37,6 +37,14 @@ The package manifest is a3s-use-extension.acl and is parsed by a3s-acl. A3S Use owns identity, routes, trust, activation, and lifecycle around the surfaces. It does not define JSON-RPC methods or convert surfaces implicitly. +`a3s-use-science` is the reference multi-surface extension. It remains a +separate process and package even though its source is developed in this +repository. Its Rust API, native CLI, 13 standard MCP tools, and packaged Skill +share typed source-specific operations; the host sees only the declared +`a3s/science` CLI, MCP, and Skill surfaces. This demonstrates how a first-party +toolkit can ship without expanding the reserved built-in route set or adding a +generic action envelope. + ## Hot-plug registry Extension code remains behind native process boundaries. The registry is a From 79a4dedafe2cbd362c204b0dcc37970ed03468c0 Mon Sep 17 00:00:00 2001 From: RoyLin Date: Sun, 19 Jul 2026 09:46:57 +0800 Subject: [PATCH 4/9] feat(ocr): add local PP-OCRv6 application capability --- .github/workflows/release.yml | 57 +- Cargo.lock | 739 +++++++++++++++++++++---- Cargo.toml | 6 + README.md | 66 ++- THIRD_PARTY_NOTICES.md | 54 ++ crates/ocr/Cargo.toml | 9 +- crates/ocr/README.md | 57 +- crates/ocr/skills/a3s-use-ocr/SKILL.md | 39 +- crates/ocr/src/assets.rs | 261 +++++++++ crates/ocr/src/cli.rs | 61 +- crates/ocr/src/client.rs | 650 ++++++---------------- crates/ocr/src/config.rs | 261 +++++++++ crates/ocr/src/engine.rs | 261 +++++++++ crates/ocr/src/install.rs | 671 ++++++++++++++++++++++ crates/ocr/src/lib.rs | 19 +- crates/ocr/src/mcp.rs | 10 +- crates/ocr/src/models.rs | 47 +- crates/ocr/src/postprocess.rs | 326 +++++++++++ crates/ocr/src/preprocess.rs | 192 +++++++ crates/ocr/tests/ppocr_v6_contract.rs | 19 + docs/architecture.md | 22 +- src/cli.rs | 62 ++- src/cli_tests.rs | 11 + src/ocr_builtin.rs | 6 +- 24 files changed, 3118 insertions(+), 788 deletions(-) create mode 100644 crates/ocr/src/assets.rs create mode 100644 crates/ocr/src/config.rs create mode 100644 crates/ocr/src/engine.rs create mode 100644 crates/ocr/src/install.rs create mode 100644 crates/ocr/src/postprocess.rs create mode 100644 crates/ocr/src/preprocess.rs create mode 100644 crates/ocr/tests/ppocr_v6_contract.rs diff --git a/.github/workflows/release.yml b/.github/workflows/release.yml index 0116089b..899d5886 100644 --- a/.github/workflows/release.yml +++ b/.github/workflows/release.yml @@ -87,6 +87,9 @@ jobs: install -d "${stage}/skills" "${stage}/skill-data" "${stage}/office-skills" "${stage}/ocr-skills" "${stage}/dashboard" install -m 0755 "target/${{ matrix.target }}/release/a3s-use" "${stage}/a3s-use" install -m 0755 "target/${{ matrix.target }}/release/a3s-use-browser-driver" "${stage}/a3s-use-browser-driver" + A3S_USE_OCR_HOME="${stage}/ocr-models" \ + "${stage}/a3s-use" component install ocr --json > "${RUNNER_TEMP}/ocr-model-install.json" + rm -f "${stage}/ocr-models/.install.lock" cp -R crates/browser-driver/skills/. "${stage}/skills/" cp -R crates/browser-driver/skill-data/. "${stage}/skill-data/" cp -R crates/office/skills/. "${stage}/office-skills/" @@ -108,6 +111,11 @@ jobs: New-Item -ItemType Directory -Force -Path $stage | Out-Null Copy-Item "target/${{ matrix.target }}/release/a3s-use.exe" "$stage/a3s-use.exe" Copy-Item "target/${{ matrix.target }}/release/a3s-use-browser-driver.exe" "$stage/a3s-use-browser-driver.exe" + $env:A3S_USE_OCR_HOME = "$stage/ocr-models" + & "$stage/a3s-use.exe" component install ocr --json | Out-File "$env:RUNNER_TEMP/ocr-model-install.json" + if ($LASTEXITCODE -ne 0) { throw "Failed to install the pinned PP-OCRv6 release assets" } + Remove-Item Env:A3S_USE_OCR_HOME + Remove-Item "$stage/ocr-models/.install.lock" -ErrorAction SilentlyContinue Copy-Item -Recurse "crates/browser-driver/skills" "$stage/skills" Copy-Item -Recurse "crates/browser-driver/skill-data" "$stage/skill-data" Copy-Item -Recurse "crates/office/skills" "$stage/office-skills" @@ -131,6 +139,10 @@ jobs: test -f "${install_root}/skill-data/core/SKILL.md" test -f "${install_root}/office-skills/a3s-use-office/SKILL.md" test -f "${install_root}/ocr-skills/a3s-use-ocr/SKILL.md" + test -f "${install_root}/ocr-models/PP-OCRv6_small/det/inference.onnx" + test -f "${install_root}/ocr-models/PP-OCRv6_small/det/inference.yml" + test -f "${install_root}/ocr-models/PP-OCRv6_small/rec/inference.onnx" + test -f "${install_root}/ocr-models/PP-OCRv6_small/rec/inference.yml" test -f "${install_root}/dashboard/index.html" test -f "${install_root}/LICENSE-APACHE-2.0" test -f "${install_root}/UPSTREAM.md" @@ -151,7 +163,27 @@ jobs: import json, pathlib, sys value = json.loads(pathlib.Path(sys.argv[1]).read_text()) assert value["ok"] is True - assert value["data"]["readiness"] in {"ready", "missing", "broken", "unknown"} + assert value["data"]["readiness"] == "ready" + assert value["data"]["provider"] == "pp-ocr-v6" + assert value["data"]["engine"] == "onnx-runtime" + assert value["data"]["model"] == "PP-OCRv6_small" + assert value["data"]["sendsSourceOffDevice"] is False + PY + python3 - "${RUNNER_TEMP}/ocr-white.bmp" <<'PY' + import base64, pathlib, sys + pathlib.Path(sys.argv[1]).write_bytes(base64.b64decode( + "Qk06AAAAAAAAADYAAAAoAAAAAQAAAAEAAAABABgAAAAAAAQAAAAAAAAAAAAAAAAAAAAAAAAA////AA==" + )) + PY + "${install_root}/a3s-use" ocr extract "${RUNNER_TEMP}/ocr-white.bmp" --json > "${RUNNER_TEMP}/ocr-extract.json" + python3 - "${RUNNER_TEMP}/ocr-extract.json" <<'PY' + import json, pathlib, sys + value = json.loads(pathlib.Path(sys.argv[1]).read_text()) + assert value["ok"] is True + assert value["data"]["provider"] == "pp-ocr-v6" + assert value["data"]["engine"] == "onnx-runtime" + assert value["data"]["model"] == "PP-OCRv6_small" + assert value["data"]["source"]["mediaType"] == "image/bmp" PY "${install_root}/a3s-use" capability snapshot --json > "${RUNNER_TEMP}/capabilities.json" python3 - "${RUNNER_TEMP}/capabilities.json" "${install_root}" <<'PY' @@ -206,6 +238,10 @@ jobs: "$root/skill-data/core/SKILL.md", "$root/office-skills/a3s-use-office/SKILL.md", "$root/ocr-skills/a3s-use-ocr/SKILL.md", + "$root/ocr-models/PP-OCRv6_small/det/inference.onnx", + "$root/ocr-models/PP-OCRv6_small/det/inference.yml", + "$root/ocr-models/PP-OCRv6_small/rec/inference.onnx", + "$root/ocr-models/PP-OCRv6_small/rec/inference.yml", "$root/dashboard/index.html", "$root/LICENSE-APACHE-2.0", "$root/UPSTREAM.md", @@ -222,7 +258,24 @@ jobs: throw "Packaged Office Skill smoke failed" } $ocr = (& "$root/a3s-use.exe" ocr doctor --json | ConvertFrom-Json) - if (-not $ocr.ok -or -not $ocr.data.readiness) { throw "Built-in OCR doctor smoke failed" } + if ( + -not $ocr.ok -or + $ocr.data.readiness -ne "ready" -or + $ocr.data.provider -ne "pp-ocr-v6" -or + $ocr.data.engine -ne "onnx-runtime" -or + $ocr.data.model -ne "PP-OCRv6_small" -or + $ocr.data.sendsSourceOffDevice + ) { throw "Packaged PP-OCRv6 doctor smoke failed" } + $whiteBmp = [Convert]::FromBase64String("Qk06AAAAAAAAADYAAAAoAAAAAQAAAAEAAAABABgAAAAAAAQAAAAAAAAAAAAAAAAAAAAAAAAA////AA==") + [IO.File]::WriteAllBytes("$env:RUNNER_TEMP/ocr-white.bmp", $whiteBmp) + $ocrExtract = (& "$root/a3s-use.exe" ocr extract "$env:RUNNER_TEMP/ocr-white.bmp" --json | ConvertFrom-Json) + if ( + -not $ocrExtract.ok -or + $ocrExtract.data.provider -ne "pp-ocr-v6" -or + $ocrExtract.data.engine -ne "onnx-runtime" -or + $ocrExtract.data.model -ne "PP-OCRv6_small" -or + $ocrExtract.data.source.mediaType -ne "image/bmp" + ) { throw "Packaged PP-OCRv6 extraction smoke failed" } $capabilities = (& "$root/a3s-use.exe" capability snapshot --json | ConvertFrom-Json) $ocrCapability = $capabilities.data.registry.capabilities | Where-Object id -eq "use/ocr" diff --git a/Cargo.lock b/Cargo.lock index 80e0f846..37fa4d99 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -121,15 +121,20 @@ name = "a3s-use-ocr" version = "0.1.1" dependencies = [ "a3s-use-core", - "axum", - "base64", "clap", + "clipper2", + "fs2", + "image", + "imageproc", + "ort", "reqwest", "rmcp", "schemars", "serde", "serde_json", + "serde_yaml", "sha2 0.10.9", + "tar", "tempfile", "tokio", "url", @@ -155,6 +160,22 @@ dependencies = [ "zip", ] +[[package]] +name = "ab_glyph" +version = "0.2.32" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "01c0457472c38ea5bd1c3b5ada5e368271cb550be7a4ca4a0b4634e9913f6cc2" +dependencies = [ + "ab_glyph_rasterizer", + "owned_ttf_parser", +] + +[[package]] +name = "ab_glyph_rasterizer" +version = "0.1.10" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "366ffbaa4442f4684d91e2cd7c5ea7c4ed8add41959a31447066e279e432b618" + [[package]] name = "adler2" version = "2.0.1" @@ -205,15 +226,6 @@ dependencies = [ "memchr", ] -[[package]] -name = "aligned" -version = "0.4.3" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "ee4508988c62edf04abd8d92897fca0c2995d907ce1dfeaf369dac3716a40685" -dependencies = [ - "as-slice", -] - [[package]] name = "aligned-vec" version = "0.6.4" @@ -288,6 +300,15 @@ version = "1.0.103" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "2a4385e2e34eb35d6b3efe798b9eb88096925d87726c0798709bf56d9ed84af3" +[[package]] +name = "approx" +version = "0.5.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "cab112f0a86d568ea0e627cc1d6be74a1e9cd55214684db5561995f6dad897c6" +dependencies = [ + "num-traits", +] + [[package]] name = "arbitrary" version = "1.4.2" @@ -314,15 +335,6 @@ version = "0.7.8" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "d3fb67a6e08acf24fdeccbac2cb6ac4305825bd1f117462e0e6f2f193345ad56" -[[package]] -name = "as-slice" -version = "0.2.1" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "516b6b4f0e40d50dcda9365d53964ec74560ad4284da2e7fc97122cd83174516" -dependencies = [ - "stable_deref_trait", -] - [[package]] name = "async-attributes" version = "1.1.2" @@ -522,26 +534,6 @@ version = "1.5.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "f2032f911046de80f0a198e0901378627c33f59ea0ac00e363d481118bd70a53" -[[package]] -name = "av-scenechange" -version = "0.14.1" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "0f321d77c20e19b92c39e7471cf986812cbb46659d2af674adc4331ef3f18394" -dependencies = [ - "aligned", - "anyhow", - "arg_enum_proc_macro", - "arrayvec", - "log", - "num-rational", - "num-traits", - "pastey", - "rayon", - "thiserror 2.0.18", - "v_frame", - "y4m", -] - [[package]] name = "av1-grain" version = "0.2.5" @@ -623,12 +615,24 @@ version = "0.22.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "72b3254f16251a8381aa12e40e3c4d2f0199f8c6508fbecb9d91f575e0fbb8c6" +[[package]] +name = "base64ct" +version = "1.8.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "2af50177e190e07a26ab74f8b1efbfe2ef87da2116221318cb1c2e82baf7de06" + [[package]] name = "bit_field" version = "0.10.3" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "1e4b40c7323adcfc0a41c4b88143ed58346ff65a288fc144329c5c45e05d70c6" +[[package]] +name = "bitflags" +version = "1.3.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "bef38d45163c2f1dde094a7dfd33ccf595c92905c8f8f4fdc18d06fb1037718a" + [[package]] name = "bitflags" version = "2.13.0" @@ -637,12 +641,9 @@ checksum = "b4388bee8683e3d04af747c73422af53102d2bd24d9eadb6cbc100baef4b43f8" [[package]] name = "bitstream-io" -version = "4.10.0" +version = "2.6.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "7eff00be299a18769011411c9def0d827e8f2d7bf0c3dbf53633147a8867fd1f" -dependencies = [ - "no_std_io2", -] +checksum = "6099cdc01846bc367c4e7dd630dc5966dccf36b652fae7a74e17b640411a91b2" [[package]] name = "block-buffer" @@ -677,9 +678,9 @@ dependencies = [ [[package]] name = "built" -version = "0.8.1" +version = "0.7.7" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "5c0e531d93d39c34eef561e929e8a7f86d77a5af08aac4f6d6e39976c51858e9" +checksum = "56ed6191a7e78c36abdb16ab65341eefd73d64d303fffccdbb00d51e4205967b" [[package]] name = "bumpalo" @@ -726,6 +727,16 @@ dependencies = [ "shlex", ] +[[package]] +name = "cfg-expr" +version = "0.15.8" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d067ad48b8650848b989a59a86c6c36a995d02d2bf778d45c3c5d57bc2718f02" +dependencies = [ + "smallvec 1.15.2", + "target-lexicon", +] + [[package]] name = "cfg-if" version = "1.0.4" @@ -881,6 +892,27 @@ version = "1.1.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "c8d4a3bb8b1e0c1050499d1815f5ab16d04f0959b233085fb31653fbfc9d98f9" +[[package]] +name = "clipper2" +version = "0.5.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "2e9c96871d5c50dd16c0f9df018f701fc62bd4f1b73cea8efe1985df6ab3d7e6" +dependencies = [ + "clipper2c-sys", + "libc", + "thiserror 2.0.18", +] + +[[package]] +name = "clipper2c-sys" +version = "0.1.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "6801d5a9a1d30e747025ec2afdc841d51665028481d61b5c719d173e88211ea3" +dependencies = [ + "cc", + "libc", +] + [[package]] name = "color_quant" version = "1.1.0" @@ -908,6 +940,16 @@ version = "0.10.2" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "a6ef517f0926dd24a1582492c791b6a4818a4d94e789a334894aa15b0d12f55c" +[[package]] +name = "core-foundation" +version = "0.10.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b2a6cd9ae233e7f62ba4e9353e81a88df7fc8a5987b8d445b4d90c879bd156f6" +dependencies = [ + "core-foundation-sys", + "libc", +] + [[package]] name = "core-foundation-sys" version = "0.8.7" @@ -1042,6 +1084,16 @@ version = "2.11.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "a4ae5f15dda3c708c0ade84bfee31ccab44a3da4f88015ed22f63732abe300c8" +[[package]] +name = "der" +version = "0.8.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a69dedd701da44b0536442edf09c81a64b0ab97a7a4a5e3d1971f00027cbc63d" +dependencies = [ + "pem-rfc7468", + "zeroize", +] + [[package]] name = "deranged" version = "0.5.8" @@ -1213,7 +1265,7 @@ dependencies = [ "num-complex", "pulp", "rayon-core", - "smallvec", + "smallvec 1.15.2", "zune-inflate", ] @@ -1223,12 +1275,6 @@ version = "2.4.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "9f1f227452a390804cdb637b74a86990f2a7d7ba4b7d5693aac9b4dd6defd8d6" -[[package]] -name = "fax" -version = "0.2.7" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "caf1079563223d5d59d83c85886a56e586cfd5c1a26292e971a0fa266531ac5a" - [[package]] name = "fdeflate" version = "0.3.7" @@ -1238,6 +1284,16 @@ dependencies = [ "simd-adler32", ] +[[package]] +name = "filetime" +version = "0.2.29" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "5c287a33c7f0a620c38e641e7f60827713987b3c0f26e8ddc9462cc69cf75759" +dependencies = [ + "cfg-if", + "libc", +] + [[package]] name = "find-msvc-tools" version = "0.1.9" @@ -1260,6 +1316,21 @@ version = "1.0.7" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "3f9eec918d3f24069decb9af1554cad7c880e2da24a9afd88aca000531ab82c1" +[[package]] +name = "foreign-types" +version = "0.3.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f6f339eb8adc052cd2ca78910fda869aefa38d22d5cb648e6485e4d3fc06f3b1" +dependencies = [ + "foreign-types-shared", +] + +[[package]] +name = "foreign-types-shared" +version = "0.1.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "00b0228411908ca8685dba7fc2cdd70ec9990a6e753e89b6ac91a84c40fbaf4b" + [[package]] name = "form_urlencoded" version = "1.2.2" @@ -1447,9 +1518,9 @@ dependencies = [ [[package]] name = "gif" -version = "0.14.2" +version = "0.13.3" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "ee8cfcc411d9adbbaba82fb72661cc1bcca13e8bba98b364e62b2dba8f960159" +checksum = "4ae047235e33e2829703574b54fdec96bfbad892062d97fed2f76022287de61b" dependencies = [ "color_quant", "weezl", @@ -1596,7 +1667,7 @@ dependencies = [ "httpdate", "itoa", "pin-project-lite", - "smallvec", + "smallvec 1.15.2", "tokio", "want", ] @@ -1701,7 +1772,7 @@ dependencies = [ "icu_normalizer_data", "icu_properties", "icu_provider", - "smallvec", + "smallvec 1.15.2", "zerovec", ] @@ -1759,7 +1830,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "3b0875f23caa03898994f6ddc501886a45c7d3d62d04d2d90788d47be1b1e4de" dependencies = [ "idna_adapter", - "smallvec", + "smallvec 1.15.2", "utf8_iter", ] @@ -1775,9 +1846,9 @@ dependencies = [ [[package]] name = "image" -version = "0.25.10" +version = "0.25.5" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "85ab80394333c02fe689eaf900ab500fbd0c2213da414687ebf995a65d5a6104" +checksum = "cd6f44aed642f18953a158afeb30206f4d50da59fbc66ecb53c66488de73563b" dependencies = [ "bytemuck", "byteorder-lite", @@ -1785,7 +1856,6 @@ dependencies = [ "exr", "gif", "image-webp", - "moxcms", "num-traits", "png", "qoi", @@ -1807,6 +1877,23 @@ dependencies = [ "quick-error", ] +[[package]] +name = "imageproc" +version = "0.25.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "2393fb7808960751a52e8a154f67e7dd3f8a2ef9bd80d1553078a7b4e8ed3f0d" +dependencies = [ + "ab_glyph", + "approx", + "getrandom 0.2.17", + "image", + "itertools", + "nalgebra", + "num", + "rand 0.8.7", + "rand_distr", +] + [[package]] name = "imgref" version = "1.12.2" @@ -1857,9 +1944,9 @@ checksum = "a6cb138bb79a146c1bd460005623e142ef0181e3d0219cb493e02f7d08a35695" [[package]] name = "itertools" -version = "0.14.0" +version = "0.12.1" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "2b192c782037fadd9cfa75548310488aabdbf3d2da73885b31bd0abd03351285" +checksum = "ba291022dbbd398a455acf126c1e341954079855bc60dfdda641363bd6922569" dependencies = [ "either", ] @@ -1880,6 +1967,12 @@ dependencies = [ "libc", ] +[[package]] +name = "jpeg-decoder" +version = "0.3.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "00810f1d8b74be64b13dbf3db89ac67740615d6c891f0e7b6179326533011a07" + [[package]] name = "js-sys" version = "0.3.103" @@ -1985,6 +2078,16 @@ version = "0.8.4" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "47e1ffaa40ddd1f3ed91f717a33c8c0ee23fff369e3aa8772b9605cc1d22f4c3" +[[package]] +name = "matrixmultiply" +version = "0.3.11" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "3f607c237553f086e7043417a51df26b2eb899d3caff94e6a67592ff992fedc7" +dependencies = [ + "autocfg", + "rawpointer", +] + [[package]] name = "maybe-rayon" version = "0.1.1" @@ -2039,30 +2142,58 @@ dependencies = [ ] [[package]] -name = "moxcms" -version = "0.8.1" +name = "nalgebra" +version = "0.32.6" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "bb85c154ba489f01b25c0d36ae69a87e4a1c73a72631fc6c0eb6dde34a73e44b" +checksum = "7b5c17de023a86f59ed79891b2e5d5a94c705dbe904a5b5c9c952ea6221b03e4" dependencies = [ + "approx", + "matrixmultiply", + "num-complex", + "num-rational", "num-traits", - "pxfm", + "simba", + "typenum", ] [[package]] -name = "new_debug_unreachable" -version = "1.0.6" +name = "native-tls" +version = "0.2.18" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "650eef8c711430f1a879fdd01d4745a7deea475becfb90269c06775983bbf086" +checksum = "465500e14ea162429d264d44189adc38b199b62b1c21eea9f69e4b73cb03bbf2" +dependencies = [ + "libc", + "log", + "openssl", + "openssl-probe", + "openssl-sys", + "schannel", + "security-framework", + "security-framework-sys", + "tempfile", +] [[package]] -name = "no_std_io2" -version = "0.9.4" +name = "ndarray" +version = "0.16.1" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "418abd1b6d34fbf6cae440dc874771b0525a604428704c76e48b29a5e67b8003" +checksum = "882ed72dce9365842bf196bdeedf5055305f11fc8c03dee7bb0194a6cad34841" dependencies = [ - "memchr", + "matrixmultiply", + "num-complex", + "num-integer", + "num-traits", + "portable-atomic", + "portable-atomic-util", + "rawpointer", ] +[[package]] +name = "new_debug_unreachable" +version = "1.0.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "650eef8c711430f1a879fdd01d4745a7deea475becfb90269c06775983bbf086" + [[package]] name = "nom" version = "8.0.0" @@ -2078,6 +2209,20 @@ version = "0.3.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "0676bb32a98c1a483ce53e500a81ad9c3d5b3f7c920c28c24e9cb0980d0b5bc8" +[[package]] +name = "num" +version = "0.4.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "35bd024e8b2ff75562e5f34e7f4905839deb4b22955ef5e73d2fea1b9813cb23" +dependencies = [ + "num-bigint", + "num-complex", + "num-integer", + "num-iter", + "num-rational", + "num-traits", +] + [[package]] name = "num-bigint" version = "0.4.8" @@ -2124,6 +2269,16 @@ dependencies = [ "num-traits", ] +[[package]] +name = "num-iter" +version = "0.1.46" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c92800bd69a1eac91786bcfe9da64a897eb72911b8dc3095decbd07429e8048b" +dependencies = [ + "num-integer", + "num-traits", +] + [[package]] name = "num-rational" version = "0.4.2" @@ -2142,6 +2297,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "071dfc062690e90b734c0b2273ce72ad0ffa95f0c74596bc250dcfd960262841" dependencies = [ "autocfg", + "libm", ] [[package]] @@ -2162,12 +2318,89 @@ version = "0.3.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "c08d65885ee38876c4f86fa503fb49d7b507c2b62552df7c70b2fce627e06381" +[[package]] +name = "openssl" +version = "0.10.81" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "77823a27f0babb03091cb9ed9ef80af3b39dbc82f97e8fa530374b7dafd87a45" +dependencies = [ + "bitflags 2.13.0", + "cfg-if", + "foreign-types", + "libc", + "openssl-macros", + "openssl-sys", +] + +[[package]] +name = "openssl-macros" +version = "0.1.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a948666b637a0f465e8564c73e89d4dde00d72d4d473cc972f390fc3dcee7d9c" +dependencies = [ + "proc-macro2", + "quote", + "syn 2.0.118", +] + +[[package]] +name = "openssl-probe" +version = "0.2.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "7c87def4c32ab89d880effc9e097653c8da5d6ef28e6b539d313baaacfbafcbe" + +[[package]] +name = "openssl-sys" +version = "0.9.117" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b47e7e6bb2c38cd930d25a23b40fa52e068c10e85f3e03a7f5ba5aaca5713695" +dependencies = [ + "cc", + "libc", + "pkg-config", + "vcpkg", +] + [[package]] name = "option-ext" version = "0.2.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "04744f49eae99ab78e0d5c0b603ab218f515ea8cfe5a456d7629ad883a3b6e7d" +[[package]] +name = "ort" +version = "2.0.0-rc.10" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1fa7e49bd669d32d7bc2a15ec540a527e7764aec722a45467814005725bcd721" +dependencies = [ + "ndarray", + "ort-sys", + "smallvec 2.0.0-alpha.10", + "tracing", +] + +[[package]] +name = "ort-sys" +version = "2.0.0-rc.10" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e2aba9f5c7c479925205799216e7e5d07cc1d4fa76ea8058c60a9a30f6a4e890" +dependencies = [ + "flate2", + "pkg-config", + "sha2 0.10.9", + "tar", + "ureq", +] + +[[package]] +name = "owned_ttf_parser" +version = "0.25.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "36820e9051aca1014ddc75770aab4d68bc1e9e632f0f5627c4086bc216fb583b" +dependencies = [ + "ttf-parser", +] + [[package]] name = "parking" version = "2.2.1" @@ -2181,10 +2414,13 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "57c0d7b74b563b49d38dae00a0c37d4d6de9b432382b2892f0574ddcae73fd0a" [[package]] -name = "pastey" -version = "0.1.1" +name = "pem-rfc7468" +version = "1.0.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "35fb2e5f958ec131621fdd531e9fc186ed768cbe395337403ae56c17a74c68ec" +checksum = "a6305423e0e7738146434843d1694d621cce767262b2a86910beab705e4493d9" +dependencies = [ + "base64ct", +] [[package]] name = "percent-encoding" @@ -2215,13 +2451,19 @@ dependencies = [ "futures-io", ] +[[package]] +name = "pkg-config" +version = "0.3.33" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "19f132c84eca552bf34cab8ec81f1c1dcc229b811638f9d283dceabe58c5569e" + [[package]] name = "png" -version = "0.18.1" +version = "0.17.16" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "60769b8b31b2a9f263dae2776c37b1b28ae246943cf719eb6946a1db05128a61" +checksum = "82151a2fc869e011c153adc57cf2789ccb8d9906ce52c0b39a6b5697749d7526" dependencies = [ - "bitflags", + "bitflags 1.3.2", "crc32fast", "fdeflate", "flate2", @@ -2254,6 +2496,21 @@ dependencies = [ "universal-hash", ] +[[package]] +name = "portable-atomic" +version = "1.14.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "3d20d5497ef88037a52ff98267d066e7f11fcc5e99bbfbd58a42336193aacec3" + +[[package]] +name = "portable-atomic-util" +version = "0.2.7" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c2a106d1259c23fac8e543272398ae0e3c0b8d33c88ed73d0cc71b0f1d902618" +dependencies = [ + "portable-atomic", +] + [[package]] name = "potential_utf" version = "0.1.5" @@ -2329,12 +2586,6 @@ version = "0.1.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "1d8f70e07b9c3962945a74e59ca1c511bba65b6419468acc217c457d93f3c740" -[[package]] -name = "pxfm" -version = "0.1.30" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "d55d956fa96f5ec02be2e13af0e20391a5aa83d6a074e3ad368959d0fab299ea" - [[package]] name = "qoi" version = "0.4.1" @@ -2512,6 +2763,16 @@ version = "0.10.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "63b8176103e19a2643978565ca18b50549f6101881c443590420e4dc998a3c69" +[[package]] +name = "rand_distr" +version = "0.4.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "32cb0b9bc82b0a0876c2dd994a7e7a2683d3e7390ca40e6886785ef0c7e3ee31" +dependencies = [ + "num-traits", + "rand 0.8.7", +] + [[package]] name = "rand_pcg" version = "0.10.2" @@ -2523,15 +2784,13 @@ dependencies = [ [[package]] name = "rav1e" -version = "0.8.1" +version = "0.7.1" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "43b6dd56e85d9483277cde964fd1bdb0428de4fec5ebba7540995639a21cb32b" +checksum = "cd87ce80a7665b1cce111f8a16c1f3929f6547ce91ade6addf4ec86a8dda5ce9" dependencies = [ - "aligned-vec", "arbitrary", "arg_enum_proc_macro", "arrayvec", - "av-scenechange", "av1-grain", "bitstream-io", "built", @@ -2546,21 +2805,23 @@ dependencies = [ "noop_proc_macro", "num-derive", "num-traits", + "once_cell", "paste", "profiling", - "rand 0.9.5", - "rand_chacha 0.9.0", + "rand 0.8.7", + "rand_chacha 0.3.1", "simd_helpers", - "thiserror 2.0.18", + "system-deps", + "thiserror 1.0.69", "v_frame", "wasm-bindgen", ] [[package]] name = "ravif" -version = "0.13.0" +version = "0.11.20" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "e52310197d971b0f5be7fe6b57530dcd27beb35c1b013f29d66c1ad73fbbcc45" +checksum = "5825c26fddd16ab9f515930d49028a630efec172e903483c94796cfe31893e6b" dependencies = [ "avif-serialize", "imgref", @@ -2577,9 +2838,15 @@ version = "11.6.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "498cd0dc59d73224351ee52a95fee0f1a617a2eae0e7d9d720cc622c73a54186" dependencies = [ - "bitflags", + "bitflags 2.13.0", ] +[[package]] +name = "rawpointer" +version = "0.2.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "60a357793950651c4ed0f3f52338f53b2f809f32d83a07f72909fa13e4c6c1e3" + [[package]] name = "rayon" version = "1.12.0" @@ -2818,7 +3085,7 @@ version = "0.38.44" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "fdb5bc1ae2baa591800df16c9ca78619bf65c0488b41b96ccec5d11220d8c154" dependencies = [ - "bitflags", + "bitflags 2.13.0", "errno", "libc", "linux-raw-sys 0.4.15", @@ -2831,7 +3098,7 @@ version = "1.1.4" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "b6fe4565b9518b83ef4f91bb47ce29620ca828bd32cb7e408f0062e9930ba190" dependencies = [ - "bitflags", + "bitflags 2.13.0", "errno", "libc", "linux-raw-sys 0.12.1", @@ -2885,6 +3152,15 @@ version = "1.0.23" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "9774ba4a74de5f7b1c1451ed6cd5285a32eddb5cccb8cc655a4e50009e06477f" +[[package]] +name = "safe_arch" +version = "0.7.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "96b02de82ddbe1b636e6170c21be622223aea188ef2e139be0a5b219ec215323" +dependencies = [ + "bytemuck", +] + [[package]] name = "same-file" version = "1.0.6" @@ -2894,6 +3170,15 @@ dependencies = [ "winapi-util", ] +[[package]] +name = "schannel" +version = "0.1.29" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "91c1b7e4904c873ef0710c1f407dde2e6287de2bebc1bbbf7d430bb7cbffd939" +dependencies = [ + "windows-sys 0.61.2", +] + [[package]] name = "schemars" version = "1.2.1" @@ -2920,6 +3205,29 @@ dependencies = [ "syn 2.0.118", ] +[[package]] +name = "security-framework" +version = "3.7.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b7f4bc775c73d9a02cde8bf7b2ec4c9d12743edf609006c7facc23998404cd1d" +dependencies = [ + "bitflags 2.13.0", + "core-foundation", + "core-foundation-sys", + "libc", + "security-framework-sys", +] + +[[package]] +name = "security-framework-sys" +version = "2.17.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "6ce2691df843ecc5d231c0b14ece2acc3efb62c0a398c7e1d875f3983ce020e3" +dependencies = [ + "core-foundation-sys", + "libc", +] + [[package]] name = "semver" version = "1.0.28" @@ -2991,6 +3299,15 @@ dependencies = [ "serde_core", ] +[[package]] +name = "serde_spanned" +version = "0.6.9" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "bf41e0cfaf7226dca15e8197172c295a782857fcb97fad1808a166870dee75a3" +dependencies = [ + "serde", +] + [[package]] name = "serde_urlencoded" version = "0.7.1" @@ -3003,6 +3320,19 @@ dependencies = [ "serde", ] +[[package]] +name = "serde_yaml" +version = "0.9.34+deprecated" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "6a8b1a1a2ebf674015cc02edccce75287f1a0130d394307b36743c2f5d504b47" +dependencies = [ + "indexmap", + "itoa", + "ryu", + "serde", + "unsafe-libyaml", +] + [[package]] name = "sha1" version = "0.10.7" @@ -3052,6 +3382,19 @@ dependencies = [ "libc", ] +[[package]] +name = "simba" +version = "0.8.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "061507c94fc6ab4ba1c9a0305018408e312e17c041eb63bef8aa726fa33aceae" +dependencies = [ + "approx", + "num-complex", + "num-traits", + "paste", + "wide", +] + [[package]] name = "simd-adler32" version = "0.3.9" @@ -3085,6 +3428,12 @@ version = "1.15.2" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "8ed6a63f02c8539c91a8685a86f4099661ba3da017932f6ebbea6de3f0fa7c90" +[[package]] +name = "smallvec" +version = "2.0.0-alpha.10" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "51d44cfb396c3caf6fbfd0ab422af02631b69ddd96d2eff0b0f0724f9024051b" + [[package]] name = "socket2" version = "0.6.5" @@ -3095,6 +3444,17 @@ dependencies = [ "windows-sys 0.61.2", ] +[[package]] +name = "socks" +version = "0.3.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f0c3dbbd9ae980613c6dd8e28a9407b50509d3803b57624d5dfe8315218cd58b" +dependencies = [ + "byteorder", + "libc", + "winapi", +] + [[package]] name = "sse-stream" version = "0.2.4" @@ -3168,6 +3528,36 @@ dependencies = [ "syn 2.0.118", ] +[[package]] +name = "system-deps" +version = "6.2.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a3e535eb8dded36d55ec13eddacd30dec501792ff23a0b1682c38601b8cf2349" +dependencies = [ + "cfg-expr", + "heck 0.5.0", + "pkg-config", + "toml", + "version-compare", +] + +[[package]] +name = "tar" +version = "0.4.46" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "3f6221d9a6003c78398e3b239969f352578258df48c8eb051caadae0015bc840" +dependencies = [ + "filetime", + "libc", + "xattr", +] + +[[package]] +name = "target-lexicon" +version = "0.12.16" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "61c41af27dd6d1e27b1b16b489db798443478cef1f06a660c96db617ba5de3b1" + [[package]] name = "tempfile" version = "3.27.0" @@ -3223,16 +3613,13 @@ dependencies = [ [[package]] name = "tiff" -version = "0.11.3" +version = "0.9.1" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "b63feaf3343d35b6ca4d50483f94843803b0f51634937cc2ec519fc32232bc52" +checksum = "ba1310fcea54c6a9a4fd1aad794ecc02c31682f6bfbecdf460bf19533eed1e3e" dependencies = [ - "fax", "flate2", - "half", - "quick-error", + "jpeg-decoder", "weezl", - "zune-jpeg", ] [[package]] @@ -3367,6 +3754,40 @@ dependencies = [ "tokio", ] +[[package]] +name = "toml" +version = "0.8.23" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "dc1beb996b9d83529a9e75c17a1686767d148d70663143c7854d8b4a09ced362" +dependencies = [ + "serde", + "serde_spanned", + "toml_datetime", + "toml_edit", +] + +[[package]] +name = "toml_datetime" +version = "0.6.11" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "22cddaf88f4fbc13c51aebbf5f8eceb5c7c5a9da2ac40a13519eb5b0a0e8f11c" +dependencies = [ + "serde", +] + +[[package]] +name = "toml_edit" +version = "0.22.27" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "41fe8c660ae4257887cf66394862d21dbca4a6ddd26f04a3560410406a2f819a" +dependencies = [ + "indexmap", + "serde", + "serde_spanned", + "toml_datetime", + "winnow", +] + [[package]] name = "tower" version = "0.5.3" @@ -3389,7 +3810,7 @@ version = "0.6.11" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "4cfcf7e2740e6fc6d4d688b4ef00650406bb94adf4731e43c096c3a19fe40840" dependencies = [ - "bitflags", + "bitflags 2.13.0", "bytes", "futures-util", "http", @@ -3451,6 +3872,12 @@ version = "0.2.5" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "e421abadd41a4225275504ea4d6566923418b7f05506fbc9c0fe86ba7396114b" +[[package]] +name = "ttf-parser" +version = "0.25.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d2df906b07856748fa3f6e0ad0cbaa047052d4a7dd609e231c4f72cee8c36f31" + [[package]] name = "tungstenite" version = "0.23.0" @@ -3517,12 +3944,48 @@ dependencies = [ "subtle", ] +[[package]] +name = "unsafe-libyaml" +version = "0.2.11" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "673aac59facbab8a9007c7f6108d11f63b603f7cabff99fabf650fea5c32b861" + [[package]] name = "untrusted" version = "0.9.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "8ecb6da28b8a351d773b68d5825ac39017e680750f980f3a1a85cd8dd28a47c1" +[[package]] +name = "ureq" +version = "3.3.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "dea7109cdcd5864d4eeb1b58a1648dc9bf520360d7af16ec26d0a9354bafcfc0" +dependencies = [ + "base64", + "der", + "log", + "native-tls", + "percent-encoding", + "rustls-pki-types", + "socks", + "ureq-proto", + "utf8-zero", + "webpki-root-certs", +] + +[[package]] +name = "ureq-proto" +version = "0.6.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e994ba84b0bd1b1b0cf92878b7ef898a5c1760108fe7b6010327e274917a808c" +dependencies = [ + "base64", + "http", + "httparse", + "log", +] + [[package]] name = "url" version = "2.5.8" @@ -3548,6 +4011,12 @@ version = "0.7.6" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "09cc8ee72d2a9becf2f2febe0205bbed8fc6615b7cb429ad062dc7b7ddd036a9" +[[package]] +name = "utf8-zero" +version = "0.8.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b8c0a043c9540bae7c578c88f91dda8bd82e59ae27c21baca69c8b191aaf5a6e" + [[package]] name = "utf8_iter" version = "1.0.4" @@ -3588,6 +4057,18 @@ version = "1.13.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "5dd4ec1eb1d240636e354a30110a1dfcb37047169a4d9bd6d9d3469df574b5c4" +[[package]] +name = "vcpkg" +version = "0.2.15" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "accd4ea62f7bb7a82fe23066fb0957d48ef677f6eeb8215f372f52e48bb32426" + +[[package]] +name = "version-compare" +version = "0.2.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "03c2856837ef78f57382f06b2b8563a2f512f7185d732608fd9176cb3b8edf0e" + [[package]] name = "version_check" version = "0.9.5" @@ -3716,6 +4197,15 @@ dependencies = [ "wasm-bindgen", ] +[[package]] +name = "webpki-root-certs" +version = "1.0.9" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b96554aa2acc8ccdb7e1c9a58a7a68dd5d13bccc69cd124cb09406db612a1c9b" +dependencies = [ + "rustls-pki-types", +] + [[package]] name = "webpki-roots" version = "0.26.11" @@ -3764,6 +4254,16 @@ dependencies = [ "winsafe", ] +[[package]] +name = "wide" +version = "0.7.33" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0ce5da8ecb62bcd8ec8b7ea19f69a51275e91299be594ea5cc6ef7819e16cd03" +dependencies = [ + "bytemuck", + "safe_arch", +] + [[package]] name = "winapi" version = "0.3.9" @@ -4011,6 +4511,15 @@ version = "0.52.6" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "589f6da84c646204747d1270a2a5661ea66ed1cced2631d546fdfb155959f9ec" +[[package]] +name = "winnow" +version = "0.7.15" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "df79d97927682d2fd8adb29682d1140b343be4ac0f08fd68b7765d9c059d3945" +dependencies = [ + "memchr", +] + [[package]] name = "winreg" version = "0.52.0" @@ -4040,10 +4549,14 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "1ffae5123b2d3fc086436f8834ae3ab053a283cfac8fe0a0b8eaae044768a4c4" [[package]] -name = "y4m" -version = "0.8.0" +name = "xattr" +version = "1.6.1" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "7a5a4b21e1a62b67a2970e6831bc091d7b87e119e7f9791aef9702e3bef04448" +checksum = "32e45ad4206f6d2479085147f02bc2ef834ac85886624a23575ae137c8aa8156" +dependencies = [ + "libc", + "rustix 1.1.4", +] [[package]] name = "yoke" @@ -4185,9 +4698,9 @@ dependencies = [ [[package]] name = "zune-core" -version = "0.5.1" +version = "0.4.12" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "cb8a0807f7c01457d0379ba880ba6322660448ddebc890ce29bb64da71fb40f9" +checksum = "3f423a2c17029964870cfaabb1f13dfab7d092a62a29a89264f4d36990ca414a" [[package]] name = "zune-inflate" @@ -4200,9 +4713,9 @@ dependencies = [ [[package]] name = "zune-jpeg" -version = "0.5.15" +version = "0.4.21" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "27bc9d5b815bc103f142aa054f561d9187d191692ec7c2d1e2b4737f8dbd7296" +checksum = "29ce2c8a9384ad323cf564b67da86e21d3cfdff87908bc1223ed5c99bc792713" dependencies = [ "zune-core", ] diff --git a/Cargo.toml b/Cargo.toml index 8c17bb9f..614fdd9a 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -23,9 +23,13 @@ async-trait = "0.1" axum = "0.8" base64 = "0.22" clap = { version = "4", features = ["derive"] } +clipper2 = { version = "=0.5.3", default-features = false } fs2 = "0.4" futures-util = "0.3" getrandom = "0.3" +image = { version = "=0.25.5", default-features = false, features = ["bmp", "gif", "jpeg", "png", "tiff", "webp"] } +imageproc = { version = "=0.25.0", default-features = false } +ort = { version = "=2.0.0-rc.10", default-features = false, features = ["copy-dylibs", "download-binaries", "ndarray", "std"] } reqwest = { version = "0.12", default-features = false, features = ["rustls-tls", "stream"] } quick-xml = "0.38" regex = "1" @@ -33,7 +37,9 @@ rmcp = { version = "=0.8.5", default-features = false, features = ["base64", "cl schemars = "1.2" serde = { version = "1", features = ["derive"] } serde_json = "1" +serde_yaml = "0.9" sha2 = "0.10" +tar = "0.4" thiserror = "2" tempfile = "3" tokio = { version = "1", features = ["fs", "io-util", "macros", "net", "rt-multi-thread", "process", "sync", "time"] } diff --git a/README.md b/README.md index 5b4916fa..41e69e45 100644 --- a/README.md +++ b/README.md @@ -100,9 +100,9 @@ a3s use mcp serve office-compat # Legacy alias: a3s use mcp serve office -# Built-in OCR; provider readiness remains explicit. +# Built-in local PP-OCRv6. a3s use ocr doctor --json -a3s use ocr extract ./scan.png --language eng --json +a3s use ocr extract ./scan.png --json a3s use mcp serve ocr ``` @@ -137,9 +137,9 @@ Every domain argument accepted by `a3s use ...` can also be passed directly to safe Word, Spreadsheet, Presentation, native MCP, and compatibility workflows - **External Domains**: Install process-isolated packages that expose any useful combination of CLI, MCP, and Skill surfaces -- **First-Party OCR Domain**: Extract text and bounded layout evidence with - a local Tesseract provider or an explicitly configured vision endpoint, - without silently installing a provider or hiding remote image transfer +- **First-Party OCR Domain**: Run pinned PP-OCRv6 detection and recognition + models locally through ONNX Runtime, with source digests and bounded layout + evidence - **Hot-Plug Discovery**: Publish immutable generation/revision snapshots so a resident host can add, replace, or remove live capabilities without restarting - **Content-Bound Skills**: Project an absolute package path and lowercase @@ -160,7 +160,7 @@ Every domain argument accepted by `a3s use ...` can also be passed directly to | Browser | Built in | Full Browser vocabulary | A3S Use standard MCP server | Six packaged Browser Skills | A3S Use | | Office | Built in | Stable Office vocabulary | Typed native preview plus OfficeCLI compatibility server | Packaged `a3s-use-office` Skill | A3S Use native engine; OfficeCLI compatibility in 0.1.x | | Box | Reserved built-in route | Native A3S Box vocabulary | — | — | Umbrella A3S CLI | -| OCR | Built in | Doctor and typed image extraction | `ocr_doctor` and `ocr_extract` | One provider-safe OCR Skill | A3S Use process and explicitly configured provider | +| OCR | Built in | Doctor and typed image extraction | `ocr_doctor` and `ocr_extract` | One local PP-OCRv6 Skill | A3S Use process with ONNX Runtime | | External domain | Installed extension | Optional native executable | Optional standard MCP server | Optional `SKILL.md` | Extension package plus A3S Use lifecycle | The Box route is component-backed. The umbrella CLI resolves its authoritative @@ -175,7 +175,7 @@ Default features are `browser`, `office`, `ocr`, `extensions`, and `mcp`. | --- | --- | | `browser` | Typed Browser library, stateless rendering, and full Browser driver delegation | | `office` | Typed Office contracts, native OOXML read engine, and temporary OfficeCLI compatibility | -| `ocr` | Built-in typed OCR CLI/MCP with local Tesseract and explicit vision providers | +| `ocr` | Built-in typed PP-OCRv6 CLI/MCP with local ONNX inference | | `extensions` | ACL manifests, package receipts, hot-plug registry, and external CLI/MCP/Skill routes | | `mcp` | Standard MCP servers plus the managed Browser Streamable HTTP lifecycle | | `lightpanda` | Explicit opt-in Lightpanda provider support in addition to Chrome | @@ -192,7 +192,7 @@ A compiled command surface is not proof that its provider is installed. Use | `a3s-use-browser-driver` | Complete interactive Browser CLI, MCP tools, Skills, Dashboard, and compatibility runtime | | `a3s-use-office` | Native OOXML foundation, typed Office operations, and compatibility lifecycle | | `a3s-use-extension` | A3S ACL manifest model, package registry, leases, and native surface descriptors | -| `a3s-use-ocr` | Typed local/vision OCR providers, CLI, MCP tools, and release-packaged Skill assets | +| `a3s-use-ocr` | Local PP-OCRv6 engine, CLI, MCP tools, pinned models, and release-packaged Skill assets | | `a3s-use` | Facade library, standalone CLI host, capability projection, and MCP entry points | ## Quick Start @@ -212,9 +212,9 @@ a3s use doctor --json Prebuilt archives are also published on [GitHub Releases](https://github.com/A3S-Lab/Use/releases). A complete archive contains `a3s-use`, its sibling `a3s-use-browser-driver`, Browser Skills, the -first-party Office and OCR Skills, the Dashboard, and license/provenance -notices. Keep those packaged assets together; installing only the facade binary -does not provide the complete Browser, Office, and OCR Skill surfaces. +first-party Office Skill, the Dashboard, and license/provenance notices. Keep +those packaged assets together; installing only the facade binary does not +provide the complete Browser and Office Skill surfaces. Build all binaries from source with: @@ -1610,27 +1610,33 @@ release packages its `a3s-use-ocr` Skill and exposes `ocr_doctor` plus `ocr_extract` over standard MCP, so a resident A3S Code session receives `mcp__use_ocr__*` without installing a separate extension. -OCR never installs a provider silently. `auto` prefers an explicitly configured -or discoverable Tesseract executable. Vision OCR is enabled only when its model -and endpoint configuration are present; non-loopback endpoints require HTTPS -and an API key, and the diagnostic discloses that the complete source image -leaves the device. Supported inputs are bounded local PNG, JPEG, WebP, GIF, +OCR has one backend: the pinned `PP-OCRv6_small` detection and recognition +models running locally through ONNX Runtime. Release archives package those +models; `a3s install use/ocr` explicitly installs or repairs the same pinned +bundle when needed. Supported inputs are bounded local PNG, JPEG, WebP, GIF, BMP, and TIFF files. The result binds the canonical source path, media type, -byte length, and SHA-256 alongside text and any available -confidence/bounding-box evidence. +byte length, and SHA-256 alongside text, recognition/detection confidence, +polygons, and bounding boxes. + +The pipeline decodes and normalizes the image, runs +`PP-OCRv6_small_det`, applies DB post-processing and reading-order sorting, +perspective-rectifies and rotates text crops, runs batched +`PP-OCRv6_small_rec`, and applies CTC decoding. It does not require Python or +PaddlePaddle, call a remote OCR API, or transfer source bytes off the device. ```bash a3s use ocr doctor --json -a3s use ocr extract ./scan.png --language eng --json +a3s use ocr extract ./scan.png --json a3s use mcp serve ocr ``` -A3S Code may first-use install the verified parent Use release. OCR provider -selection remains explicit, and remote vision extraction still escalates to -the parent TUI before source bytes leave the device. +A3S Code may first-use install the verified parent Use release. A missing or +damaged managed model bundle is repaired explicitly with +`a3s install use/ocr`; the Code `use` worker never installs it implicitly. + +See the [OCR crate](crates/ocr/README.md) for model resolution, the inference +workflow, and input boundaries. -See the [OCR crate](crates/ocr/README.md) for configuration and provider -boundaries. ## External Extensions External Use domains stay behind process boundaries. A package contains an @@ -1760,7 +1766,7 @@ crash, and in-flight calls retain the exact package generation they accepted. ┌──────────┬──────────┬──────────┬──────────────┐ │ │ │ │ │ Browser Office OCR extension registry - typed + driver OOXML local/vision CLI / MCP / Skill + typed + driver OOXML PP-OCRv6 ONNX CLI / MCP / Skill + 0.1 compat │ │ │ │ └──────── capability snapshot/watch ───────────► A3S Code @@ -1773,11 +1779,11 @@ crash, and in-flight calls retain the exact package generation they accepted. The dependency arrows are intentional. Search links only the Browser contract, so rendering does not require `a3s-use`, MCP, or a resident process. Office is an in-process typed engine with a temporary 0.1.x compatibility process. OCR -uses an explicitly present local Tesseract executable or an explicitly -configured vision provider; it never installs either silently. External -domains retain their process boundaries. A3S Code consumes the read-only -projection and connects standard MCP/Skill surfaces; bounded provider -installation requests still require the parent TUI's authority. +runs the pinned PP-OCRv6 models locally through ONNX Runtime; model installation +is an explicit component operation. External domains retain their process +boundaries. A3S Code consumes the read-only projection and connects standard +MCP/Skill surfaces; bounded component installation requests still require the +parent TUI's authority. Source is split between the facade under `src/` and focused workspace crates under `crates/`. See [Architecture](docs/architecture.md) for package leases, diff --git a/THIRD_PARTY_NOTICES.md b/THIRD_PARTY_NOTICES.md index 596fd092..8040153b 100644 --- a/THIRD_PARTY_NOTICES.md +++ b/THIRD_PARTY_NOTICES.md @@ -11,3 +11,57 @@ license text is distributed as `LICENSE-APACHE-2.0`, and detailed provenance is distributed as `UPSTREAM.md`. Upstream repository: + +## PaddlePaddle/PaddleOCR PP-OCRv6 Models + +A3S Use release archives redistribute the official +`PP-OCRv6_small_det` and `PP-OCRv6_small_rec` ONNX inference model bundles +published by PaddlePaddle/PaddleOCR. The installer pins the upstream archive +URLs, byte sizes, and SHA-256 digests and does not modify the model weights. + +PaddleOCR is licensed under the Apache License, Version 2.0. + +Upstream repository: + +Model collection: + +## Microsoft ONNX Runtime + +`a3s-use-ocr` executes the models with Microsoft ONNX Runtime 1.22.0, obtained +through the pinned `ort`/`ort-sys` Rust dependencies. ONNX Runtime is licensed +under the MIT License. + +Copyright (c) Microsoft Corporation. + +Upstream repository: + +## pykeio/ort + +The `ort` and `ort-sys` Rust crates, version `2.0.0-rc.10`, provide the native +ONNX Runtime bindings and build integration. They are available under the MIT +License or the Apache License, Version 2.0. + +Upstream repository: + +## image-rs/imageproc + +`a3s-use-ocr` uses `imageproc` version `0.25.0` for geometric image +transformations. `imageproc` is licensed under the MIT License. + +Copyright (c) 2015 PistonDevelopers. + +Upstream repository: + +## clipper2 + +`a3s-use-ocr` uses the `clipper2` Rust crate version `0.5.3` and +`clipper2c-sys` version `0.1.6` for bounded polygon offsetting during DB +post-processing. The Rust crates are available under the MIT License or the +Apache License, Version 2.0. Their bundled Clipper2 C/C++ implementation is +licensed under the Boost Software License, Version 1.0. + +Upstream repositories: + +- +- +- diff --git a/crates/ocr/Cargo.toml b/crates/ocr/Cargo.toml index 4a5f58aa..dac1f616 100644 --- a/crates/ocr/Cargo.toml +++ b/crates/ocr/Cargo.toml @@ -18,17 +18,22 @@ path = "src/main.rs" [dependencies] a3s-use-core = { version = "0.1.1", path = "../core" } -base64.workspace = true clap.workspace = true +clipper2.workspace = true +fs2.workspace = true +image.workspace = true +imageproc.workspace = true +ort.workspace = true reqwest = { workspace = true, features = ["json"] } rmcp.workspace = true schemars.workspace = true serde.workspace = true serde_json.workspace = true +serde_yaml.workspace = true sha2.workspace = true +tar.workspace = true tokio.workspace = true url.workspace = true [dev-dependencies] -axum.workspace = true tempfile.workspace = true diff --git a/crates/ocr/README.md b/crates/ocr/README.md index d61b4096..a3b6e624 100644 --- a/crates/ocr/README.md +++ b/crates/ocr/README.md @@ -2,33 +2,50 @@ `a3s-use-ocr` implements the first-party built-in OCR domain for A3S Use. A3S Code receives it as `mcp__use_ocr__*` through the release-matched Use registry, -without a separate extension install. It exposes the same typed extraction -through a native CLI and standard stdio MCP, and does not silently install an -OCR provider. +without installing a separate extension. The native CLI and standard stdio MCP +share one local PP-OCRv6 implementation. -Provider selection is explicit: +There is one OCR provider: -- `A3S_OCR_PROVIDER=auto|tesseract|vision` -- `A3S_OCR_TESSERACT_EXECUTABLE=/absolute/path/to/tesseract` -- `A3S_OCR_VISION_MODEL=` -- `A3S_OCR_VISION_BASE_URL=https://provider.example/v1/` -- `A3S_OCR_VISION_API_KEY=` -- `A3S_OCR_TIMEOUT_MS=60000` +- provider: `pp-ocr-v6` +- engine: `onnx-runtime` +- model bundle: `PP-OCRv6_small` -`auto` prefers a configured or discoverable local Tesseract executable. It uses -the vision provider only when the vision environment is configured. Remote -vision endpoints require HTTPS and an API key; loopback HTTP is allowed for a -local provider. +The release packages the pinned detection and recognition models. If the model +bundle is absent or damaged, install or repair it explicitly: -Build and exercise the domain through the Use facade: +```bash +a3s install use/ocr +a3s install use/ocr --force +``` + +`A3S_OCR_MODEL_DIR` can point development builds at an explicit model bundle. +`A3S_USE_OCR_HOME` overrides the managed model root for packaging, tests, or an +isolated installation. Neither setting selects another OCR backend. + +## Workflow + +For each bounded local image, the native engine: + +1. decodes the image and applies PP-OCRv6 BGR normalization; +2. runs `PP-OCRv6_small_det` through ONNX Runtime; +3. applies DB post-processing, polygon unclipping, and reading-order sorting; +4. perspective-rectifies each text polygon and rotates tall crops; +5. runs batched `PP-OCRv6_small_rec` inference; and +6. applies CTC decoding and returns text, recognition/detection confidence, + polygons, bounding boxes, and the source SHA-256. + +All inference stays in the local `a3s-use` process. It does not require Python +or PaddlePaddle, does not call an OCR API, and does not transfer image bytes off +the device. + +## Commands ```bash a3s use ocr doctor --json -a3s use ocr extract ./scan.png --language eng --json +a3s use ocr extract ./scan.png --json a3s use mcp serve ocr ``` -The A3S Use release packages the OCR Skill beside the facade binary. A3S Code -can first-use install that verified release and hot-plug the built-in route. -Provider setup remains explicit: local Tesseract never sends source bytes -off-device, while a configured remote vision provider requires parent HITL. +Supported inputs are bounded local PNG, JPEG, WebP, GIF, BMP, and TIFF files. +URLs and PDF rasterization are outside this crate. diff --git a/crates/ocr/skills/a3s-use-ocr/SKILL.md b/crates/ocr/skills/a3s-use-ocr/SKILL.md index 9ca0950b..0a876eaf 100644 --- a/crates/ocr/skills/a3s-use-ocr/SKILL.md +++ b/crates/ocr/skills/a3s-use-ocr/SKILL.md @@ -1,36 +1,37 @@ --- name: a3s-use-ocr -description: Extract text and layout evidence from local image files through the built-in A3S Use OCR domain. Use when an agent needs optical character recognition for a PNG, JPEG, WebP, GIF, BMP, or TIFF image and must preserve the source digest, provider disclosure, confidence, and bounding-box evidence. +description: Extract text and layout evidence from local image files through the built-in A3S Use PP-OCRv6 domain. Use when an agent needs optical character recognition for a PNG, JPEG, WebP, GIF, BMP, or TIFF image and must preserve the source digest, confidence, polygon, and bounding-box evidence. --- # A3S Use OCR Use the host-provided A3S Use surface. In an A3S Code `use` worker, call `mcp__use_ocr__ocr_doctor` and `mcp__use_ocr__ocr_extract` directly. The host -owns the MCP process; do not run a shell command, install a provider, or read the +owns the MCP process; do not run a shell command, install models, or read the file through another tool. ## Workflow 1. Call `mcp__use_ocr__ocr_doctor`. -2. Confirm which provider is ready and whether `sendsSourceOffDevice` is true. +2. Confirm that `pp-ocr-v6`, `onnx-runtime`, and `PP-OCRv6_small` are ready. 3. Call `mcp__use_ocr__ocr_extract` with the exact local image path from the - task. Supply language identifiers only when known. -4. Preserve the returned source path, media type, size, and SHA-256 in the - result. Treat text, confidence, and bounding boxes as OCR evidence, not as a - verified transcription. - -The local Tesseract provider does not send the image over the network. The -vision provider sends the complete source image and prompt to its disclosed -endpoint. Do not use a non-loopback vision provider unless the user has -authorized that data transfer. Never install, repair, or switch providers from -inside the `use` worker. + task. +4. Preserve the returned source path, media type, size, and SHA-256. Treat the + decoded text, recognition/detection confidence, polygons, and bounding boxes + as OCR evidence rather than verified source text. + +The engine runs detection, reading-order sorting, perspective crop correction, +tall-crop rotation, recognition, and CTC decoding locally. It does not require +Python or PaddlePaddle and never sends the source image off the device. If the +doctor reports missing or damaged models, return its typed error and explicit +`a3s install use/ocr` suggestion to the parent; never install or repair models +from inside the `use` worker. In a CLI-only host, equivalent commands are: ```bash a3s use ocr doctor --json -a3s use ocr extract "$IMAGE" --language eng --json +a3s use ocr extract "$IMAGE" --json ``` `a3s-use-ocr` accepts the same arguments when invoked as a standalone @@ -40,9 +41,7 @@ development binary. - Only bounded local image files are accepted. URLs and PDF rasterization are outside this domain. -- Keep the default prompt for faithful transcription. A custom vision prompt - must remain an extraction instruction; do not ask the provider to interpret - unrelated content. -- Never report vision output as calibrated confidence or layout evidence. -- Do not silently fall back from a requested provider. Report typed provider, - source, and response errors to the parent agent. +- Do not ask OCR to interpret unrelated content or present OCR output as + verified source text. +- Do not hide empty results, warnings, model readiness failures, or source + digest evidence from the parent agent. diff --git a/crates/ocr/src/assets.rs b/crates/ocr/src/assets.rs new file mode 100644 index 00000000..e60e23fa --- /dev/null +++ b/crates/ocr/src/assets.rs @@ -0,0 +1,261 @@ +use std::path::{Path, PathBuf}; + +use a3s_use_core::{UseError, UseResult}; +use serde::{Deserialize, Serialize}; + +use crate::config::{load_detection, load_recognition, MODEL_FAMILY}; + +pub(crate) const RECEIPT_FILE: &str = ".a3s-ppocr-v6.json"; + +#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)] +#[serde(rename_all = "kebab-case")] +pub enum OcrInstallSource { + Environment, + Packaged, + Managed, + Missing, +} + +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +#[serde(rename_all = "camelCase")] +pub struct OcrRuntimeStatus { + pub available: bool, + pub source: OcrInstallSource, + pub model: String, + pub model_dir: Option, + pub managed_root: Option, + pub detail: String, +} + +#[derive(Debug, Clone)] +pub(crate) struct ModelAssets { + pub(crate) root: PathBuf, + pub(crate) detection_model: PathBuf, + pub(crate) detection_config: PathBuf, + pub(crate) recognition_model: PathBuf, + pub(crate) recognition_config: PathBuf, + pub(crate) source: OcrInstallSource, +} + +pub fn ocr_status() -> OcrRuntimeStatus { + let managed_root = managed_root().ok(); + match resolve_model_assets() { + Ok(assets) => OcrRuntimeStatus { + available: true, + source: assets.source, + model: MODEL_FAMILY.to_string(), + model_dir: Some(assets.root), + managed_root, + detail: "ready".to_string(), + }, + Err(error) => OcrRuntimeStatus { + available: false, + source: error + .details + .get("source") + .and_then(serde_json::Value::as_str) + .map(source_from_name) + .unwrap_or(OcrInstallSource::Missing), + model: MODEL_FAMILY.to_string(), + model_dir: error + .details + .get("modelDir") + .and_then(serde_json::Value::as_str) + .map(PathBuf::from), + managed_root, + detail: error.message, + }, + } +} + +pub(crate) fn resolve_model_assets() -> UseResult { + if let Some(path) = std::env::var_os("A3S_OCR_MODEL_DIR") + .filter(|value| !value.is_empty()) + .map(PathBuf::from) + { + let path = absolute(path)?; + return validate_assets(&path, OcrInstallSource::Environment); + } + + let managed = managed_model_dir()?; + if path_exists(&managed)? { + return validate_assets(&managed, OcrInstallSource::Managed); + } + + if let Ok(executable) = std::env::current_exe() { + if let Some(parent) = executable.parent() { + let packaged = parent.join("ocr-models").join(MODEL_FAMILY); + if path_exists(&packaged)? { + return validate_assets(&packaged, OcrInstallSource::Packaged); + } + } + } + + Err(UseError::new( + "use.ocr.model_missing", + format!("The local {MODEL_FAMILY} model bundle is not installed."), + ) + .with_suggestion("Run 'a3s install use/ocr'.") + .with_detail("source", "missing") + .with_detail("modelDir", managed.display().to_string())) +} + +pub(crate) fn validate_assets(root: &Path, source: OcrInstallSource) -> UseResult { + let root = std::fs::canonicalize(root).map_err(|error| { + model_error( + source, + root, + format!( + "Failed to resolve the {MODEL_FAMILY} model directory '{}': {error}", + root.display() + ), + ) + })?; + let detection_model = checked_file(&root, "det/inference.onnx", 256 * 1024 * 1024, source)?; + let detection_config = checked_file(&root, "det/inference.yml", 2 * 1024 * 1024, source)?; + let recognition_model = checked_file(&root, "rec/inference.onnx", 256 * 1024 * 1024, source)?; + let recognition_config = checked_file(&root, "rec/inference.yml", 2 * 1024 * 1024, source)?; + + load_detection(&detection_config)?; + load_recognition(&recognition_config)?; + + Ok(ModelAssets { + root, + detection_model, + detection_config, + recognition_model, + recognition_config, + source, + }) +} + +pub(crate) fn managed_root() -> UseResult { + if let Some(value) = std::env::var_os("A3S_USE_OCR_HOME") { + return absolute(PathBuf::from(value)); + } + if let Some(value) = std::env::var_os("A3S_DATA_HOME") { + return Ok(absolute(PathBuf::from(value))?.join("use/ocr")); + } + if let Some(value) = std::env::var_os("XDG_DATA_HOME") { + return Ok(absolute(PathBuf::from(value))?.join("a3s/use/ocr")); + } + if let Some(home) = std::env::var_os("HOME").map(PathBuf::from) { + return Ok(absolute(home)?.join(".local/share/a3s/use/ocr")); + } + #[cfg(windows)] + if let Some(value) = std::env::var_os("LOCALAPPDATA") { + return Ok(absolute(PathBuf::from(value))?.join("a3s/use/ocr")); + } + Err(UseError::new( + "use.ocr.data_home_missing", + "Cannot determine the A3S Use OCR data directory.", + )) +} + +pub(crate) fn managed_model_dir() -> UseResult { + Ok(managed_root()?.join(MODEL_FAMILY)) +} + +fn checked_file( + root: &Path, + relative: &str, + max_bytes: u64, + source: OcrInstallSource, +) -> UseResult { + let path = root.join(relative); + let canonical = std::fs::canonicalize(&path).map_err(|error| { + model_error( + source, + root, + format!( + "Required {MODEL_FAMILY} asset '{}' is unreadable: {error}", + path.display() + ), + ) + })?; + if !canonical.starts_with(root) { + return Err(model_error( + source, + root, + format!( + "Required {MODEL_FAMILY} asset '{}' escapes its model directory.", + path.display() + ), + )); + } + let metadata = std::fs::metadata(&canonical).map_err(|error| { + model_error( + source, + root, + format!( + "Failed to inspect {MODEL_FAMILY} asset '{}': {error}", + canonical.display() + ), + ) + })?; + if !metadata.is_file() || metadata.len() == 0 || metadata.len() > max_bytes { + return Err(model_error( + source, + root, + format!( + "{MODEL_FAMILY} asset '{}' must be a non-empty regular file no larger than {max_bytes} bytes.", + canonical.display() + ), + )); + } + Ok(canonical) +} + +fn path_exists(path: &Path) -> UseResult { + match std::fs::symlink_metadata(path) { + Ok(_) => Ok(true), + Err(error) if error.kind() == std::io::ErrorKind::NotFound => Ok(false), + Err(error) => Err(UseError::new( + "use.ocr.model_unreadable", + format!( + "Failed to inspect OCR model path '{}': {error}", + path.display() + ), + )), + } +} + +fn model_error(source: OcrInstallSource, root: &Path, message: impl Into) -> UseError { + UseError::new("use.ocr.model_invalid", message) + .with_suggestion("Run 'a3s install use/ocr --force' to restore the pinned PP-OCRv6 bundle.") + .with_detail("source", source_name(source)) + .with_detail("modelDir", root.display().to_string()) +} + +fn source_name(source: OcrInstallSource) -> &'static str { + match source { + OcrInstallSource::Environment => "environment", + OcrInstallSource::Packaged => "packaged", + OcrInstallSource::Managed => "managed", + OcrInstallSource::Missing => "missing", + } +} + +fn source_from_name(value: &str) -> OcrInstallSource { + match value { + "environment" => OcrInstallSource::Environment, + "packaged" => OcrInstallSource::Packaged, + "managed" => OcrInstallSource::Managed, + _ => OcrInstallSource::Missing, + } +} + +fn absolute(path: PathBuf) -> UseResult { + if path.is_absolute() { + Ok(path) + } else { + std::env::current_dir() + .map(|directory| directory.join(path)) + .map_err(|error| { + UseError::new( + "use.ocr.path_resolution_failed", + format!("Failed to resolve OCR data path: {error}"), + ) + }) + } +} diff --git a/crates/ocr/src/cli.rs b/crates/ocr/src/cli.rs index ea76bcda..f5e3e690 100644 --- a/crates/ocr/src/cli.rs +++ b/crates/ocr/src/cli.rs @@ -2,10 +2,10 @@ use std::path::PathBuf; use a3s_use_core::{UseError, UseResult}; use clap::error::ErrorKind; -use clap::{Parser, Subcommand, ValueEnum}; +use clap::{Parser, Subcommand}; use serde::Serialize; -use crate::{OcrClient, OcrMcpServer, OcrProviderKind, OcrRequest}; +use crate::{OcrClient, OcrMcpServer, OcrRequest}; #[derive(Debug)] pub struct CommandOutput { @@ -75,24 +75,10 @@ struct Cli { #[derive(Debug, Subcommand)] enum Command { - /// Inspect provider readiness without reading an image. + /// Inspect local PP-OCRv6 readiness without reading an image. Doctor, - /// Extract text and available layout evidence from one local image. - Extract { - path: PathBuf, - /// OCR language identifier; may be repeated. - #[arg(long = "language")] - languages: Vec, - /// Tesseract page segmentation mode from 0 through 13. - #[arg(long = "psm")] - page_segmentation_mode: Option, - /// Override the configured OCR provider for this call. - #[arg(long, value_enum)] - provider: Option, - /// Vision-only extraction instruction. - #[arg(long)] - prompt: Option, - }, + /// Extract text and layout evidence from one local image. + Extract { path: PathBuf }, /// Run an extension protocol surface. Serve { /// Serve standard MCP over stdin/stdout. @@ -101,23 +87,6 @@ enum Command { }, } -#[derive(Debug, Clone, Copy, ValueEnum)] -enum ProviderArg { - Auto, - Tesseract, - Vision, -} - -impl From for OcrProviderKind { - fn from(value: ProviderArg) -> Self { - match value { - ProviderArg::Auto => Self::Auto, - ProviderArg::Tesseract => Self::Tesseract, - ProviderArg::Vision => Self::Vision, - } - } -} - pub async fn run(args: Vec) -> UseResult { let mut argv = vec!["a3s-use-ocr".to_string()]; argv.extend(args); @@ -148,23 +117,9 @@ pub async fn run(args: Vec) -> UseResult { let client = OcrClient::from_env()?; match cli.command { Command::Doctor => CommandOutput::data(client.diagnostic()), - Command::Extract { - path, - languages, - page_segmentation_mode, - provider, - prompt, - } => CommandOutput::data( - client - .extract(OcrRequest { - path, - languages, - page_segmentation_mode, - provider: provider.map(Into::into), - prompt, - }) - .await?, - ), + Command::Extract { path } => { + CommandOutput::data(client.extract(OcrRequest { path }).await?) + } Command::Serve { .. } => Err(UseError::new( "use.ocr.command_invalid", "OCR MCP command dispatch reached an invalid state.", diff --git a/crates/ocr/src/client.rs b/crates/ocr/src/client.rs index 21fb3494..9dc3e33c 100644 --- a/crates/ocr/src/client.rs +++ b/crates/ocr/src/client.rs @@ -1,216 +1,178 @@ -use std::collections::BTreeMap; -use std::path::Path; -#[cfg(all(test, unix))] -use std::path::PathBuf; -use std::process::Stdio; +use std::path::{Path, PathBuf}; +use std::sync::{Arc, Mutex}; -use a3s_use_core::{Artifact, UseError, UseResult}; -use base64::Engine; +use a3s_use_core::{Artifact, Readiness, UseError, UseResult}; use sha2::{Digest, Sha256}; use tokio::io::AsyncReadExt; -use tokio::process::Command; -use crate::models::{OcrBlock, OcrBoundingBox, OcrProviderKind, OcrRequest, OcrResult}; -use crate::provider::{Provider, ProviderConfig}; -use crate::OcrDiagnostic; +use crate::assets::{ocr_status, resolve_model_assets, OcrInstallSource}; +use crate::config::MODEL_FAMILY; +use crate::engine::{EngineBlock, PpOcrV6Engine}; +use crate::models::{ + OcrBlock, OcrBoundingBox, OcrDiagnostic, OcrPoint, OcrProviderKind, OcrRequest, OcrResult, +}; +use crate::preprocess::decode_image; const MAX_INPUT_BYTES: u64 = 32 * 1024 * 1024; -const MAX_PROVIDER_OUTPUT_BYTES: usize = 8 * 1024 * 1024; -const DEFAULT_VISION_PROMPT: &str = "Transcribe all visible text in reading order. Preserve line breaks and meaningful spacing. Return only the transcription; do not summarize, translate, or wrap it in Markdown."; +const ENGINE_NAME: &str = "onnx-runtime"; #[derive(Clone)] pub struct OcrClient { - providers: ProviderConfig, - http: reqwest::Client, + loaded: Arc>>, +} + +struct LoadedEngine { + model_dir: PathBuf, + engine: PpOcrV6Engine, } impl OcrClient { pub fn from_env() -> UseResult { - Self::from_provider_config(ProviderConfig::from_env()?) - } - - fn from_provider_config(providers: ProviderConfig) -> UseResult { - let http = reqwest::Client::builder() - .user_agent(concat!("a3s-use-ocr/", env!("CARGO_PKG_VERSION"))) - .build() - .map_err(|error| { - UseError::new( - "use.ocr.client_failed", - format!("Failed to initialize the OCR HTTP client: {error}"), - ) - })?; - Ok(Self { providers, http }) - } - - #[cfg(all(test, unix))] - pub(crate) fn with_tesseract(executable: PathBuf) -> UseResult { - Self::from_provider_config(ProviderConfig::tesseract(executable)) + Ok(Self { + loaded: Arc::new(Mutex::new(None)), + }) } pub fn diagnostic(&self) -> OcrDiagnostic { - self.providers.diagnostic() + let status = ocr_status(); + let (readiness, suggestions) = if status.available { + (Readiness::Ready, Vec::new()) + } else if status.source == OcrInstallSource::Missing { + ( + Readiness::Missing, + vec![ + "Run 'a3s install use/ocr' to install the pinned local model bundle." + .to_string(), + ], + ) + } else { + ( + Readiness::Broken, + vec![ + "Run 'a3s install use/ocr --force' to restore the pinned local model bundle." + .to_string(), + ], + ) + }; + OcrDiagnostic { + readiness, + provider: Some(OcrProviderKind::PpOcrV6), + engine: Some(ENGINE_NAME.to_string()), + model: Some(status.model), + model_dir: status.model_dir, + sends_source_off_device: false, + message: if status.available { + "Local PP-OCRv6 detection and recognition models are ready.".to_string() + } else { + status.detail + }, + suggestions, + } } pub async fn extract(&self, request: OcrRequest) -> UseResult { - validate_request(&request)?; let source = read_source(&request.path).await?; - let provider = self - .providers - .resolve(request.provider.unwrap_or(OcrProviderKind::Auto))?; - let languages = if request.languages.is_empty() { - vec!["eng".to_string()] - } else { - request.languages.clone() - }; - - let (text, blocks, warnings) = match &provider { - Provider::Tesseract { - executable, - timeout, - } => { - let output = run_tesseract( - executable, - &source.artifact.path, - &languages, - request.page_segmentation_mode, - *timeout, + let loaded = Arc::clone(&self.loaded); + tokio::task::spawn_blocking(move || { + let image = decode_image(&source.bytes)?; + let assets = resolve_model_assets()?; + let mut loaded = loaded.lock().map_err(|_| { + UseError::new( + "use.ocr.runtime_failed", + "The local PP-OCRv6 engine lock is poisoned.", ) - .await?; - let (text, blocks) = parse_tesseract_tsv(&output)?; - (text, blocks, Vec::new()) + })?; + let should_load = loaded + .as_ref() + .map(|loaded| loaded.model_dir != assets.root) + .unwrap_or(true); + if should_load { + *loaded = Some(LoadedEngine { + model_dir: assets.root.clone(), + engine: PpOcrV6Engine::load(&assets)?, + }); } - Provider::Vision { - endpoint, - api_key, - model, - timeout, - } => { - let text = self - .run_vision( - endpoint, - api_key.as_deref(), - model, - &source, - request.prompt.as_deref(), - *timeout, - ) - .await?; - let blocks = (!text.is_empty()) - .then(|| OcrBlock { - page: 1, - text: text.clone(), - confidence: None, - bounding_box: None, - }) - .into_iter() - .collect(); - ( - text, - blocks, - vec![ - "The vision provider does not return calibrated OCR confidence or bounding boxes." - .to_string(), - ], + let engine = loaded.as_mut().ok_or_else(|| { + UseError::new( + "use.ocr.runtime_failed", + "The local PP-OCRv6 engine failed to initialize.", ) - } - }; - - Ok(OcrResult { - provider: provider.kind(), - source: source.artifact, - languages, - text, - blocks, - warnings, + })?; + let blocks = engine.engine.extract(&image)?; + build_result(source.artifact, blocks) }) - } - - async fn run_vision( - &self, - endpoint: &url::Url, - api_key: Option<&str>, - model: &str, - source: &SourceImage, - prompt: Option<&str>, - timeout: std::time::Duration, - ) -> UseResult { - let encoded = base64::engine::general_purpose::STANDARD.encode(&source.bytes); - let data_url = format!("data:{};base64,{encoded}", source.artifact.media_type); - let prompt = prompt - .map(str::trim) - .filter(|value| !value.is_empty()) - .unwrap_or(DEFAULT_VISION_PROMPT); - let body = serde_json::json!({ - "model": model, - "temperature": 0, - "messages": [{ - "role": "user", - "content": [ - { "type": "text", "text": prompt }, - { - "type": "image_url", - "image_url": { - "url": data_url, - "detail": "high" - } - } - ] - }] - }); - let mut request = self - .http - .post(endpoint.clone()) - .timeout(timeout) - .json(&body); - if let Some(api_key) = api_key { - request = request.bearer_auth(api_key); - } - let response = request.send().await.map_err(|error| { - UseError::new( - "use.ocr.vision_request_failed", - format!("The vision OCR request failed: {error}"), - ) - .with_detail("endpoint", redacted_endpoint(endpoint)) - })?; - let status = response.status(); - let bytes = response.bytes().await.map_err(|error| { - UseError::new( - "use.ocr.vision_response_invalid", - format!("Failed to read the vision OCR response: {error}"), - ) - })?; - if bytes.len() > MAX_PROVIDER_OUTPUT_BYTES { - return Err(UseError::new( - "use.ocr.output_too_large", - "The vision OCR provider response exceeded 8 MiB.", - )); - } - if !status.is_success() { - let message = String::from_utf8_lossy(&bytes); - return Err(UseError::new( - "use.ocr.vision_request_failed", - format!( - "The vision OCR provider returned HTTP {status}: {}", - bounded_text(&message, 1024) - ), - ) - .with_detail("status", u64::from(status.as_u16()))); - } - let value: serde_json::Value = serde_json::from_slice(&bytes).map_err(|error| { - UseError::new( - "use.ocr.vision_response_invalid", - format!("The vision OCR provider returned invalid JSON: {error}"), - ) - })?; - let content = value.pointer("/choices/0/message/content").ok_or_else(|| { + .await + .map_err(|error| { UseError::new( - "use.ocr.vision_response_invalid", - "The vision OCR response did not contain choices[0].message.content.", + "use.ocr.runtime_failed", + format!("The local PP-OCRv6 inference task failed: {error}"), ) - })?; - let text = vision_content_text(content)?; - Ok(text.trim().to_string()) + })? + } +} + +fn build_result(source: Artifact, blocks: Vec) -> UseResult { + let blocks = blocks + .into_iter() + .map(|block| { + let [first, second, third, fourth] = block.polygon; + let polygon = [ + ocr_point(first)?, + ocr_point(second)?, + ocr_point(third)?, + ocr_point(fourth)?, + ]; + let min_x = polygon.iter().map(|point| point.x).min().unwrap_or(0); + let max_x = polygon.iter().map(|point| point.x).max().unwrap_or(0); + let min_y = polygon.iter().map(|point| point.y).min().unwrap_or(0); + let max_y = polygon.iter().map(|point| point.y).max().unwrap_or(0); + Ok(OcrBlock { + page: 1, + text: block.text, + confidence: block.confidence, + detection_confidence: block.detection_confidence, + polygon, + bounding_box: OcrBoundingBox { + x: min_x, + y: min_y, + width: max_x.saturating_sub(min_x), + height: max_y.saturating_sub(min_y), + }, + }) + }) + .collect::>>()?; + let text = blocks + .iter() + .filter(|block| !block.text.trim().is_empty()) + .map(|block| block.text.as_str()) + .collect::>() + .join("\n"); + Ok(OcrResult { + provider: OcrProviderKind::PpOcrV6, + engine: ENGINE_NAME.to_string(), + model: MODEL_FAMILY.to_string(), + source, + text, + blocks, + warnings: Vec::new(), + }) +} + +fn ocr_point(point: imageproc::point::Point) -> UseResult { + Ok(OcrPoint { + x: finite_coordinate(point.x)?, + y: finite_coordinate(point.y)?, + }) +} + +fn finite_coordinate(value: f32) -> UseResult { + if !value.is_finite() || value < 0.0 || value > u32::MAX as f32 { + return Err(UseError::new( + "use.ocr.provider_output_invalid", + "PP-OCRv6 returned an invalid polygon coordinate.", + )); } + Ok(value.round() as u32) } struct SourceImage { @@ -303,247 +265,6 @@ async fn read_source(path: &Path) -> UseResult { }) } -fn validate_request(request: &OcrRequest) -> UseResult<()> { - if request.languages.len() > 16 { - return Err(UseError::new( - "use.ocr.languages_invalid", - "At most 16 OCR language identifiers may be requested.", - )); - } - for language in &request.languages { - if language.is_empty() - || language.len() > 32 - || !language - .bytes() - .all(|byte| byte.is_ascii_alphanumeric() || matches!(byte, b'_' | b'-')) - { - return Err(UseError::new( - "use.ocr.languages_invalid", - format!("OCR language identifier '{language}' is invalid."), - )); - } - } - if request.page_segmentation_mode.is_some_and(|mode| mode > 13) { - return Err(UseError::new( - "use.ocr.page_segmentation_invalid", - "Tesseract page segmentation mode must be from 0 through 13.", - )); - } - if request - .prompt - .as_ref() - .is_some_and(|prompt| prompt.len() > 8 * 1024) - { - return Err(UseError::new( - "use.ocr.prompt_too_large", - "The vision OCR prompt must not exceed 8192 bytes.", - )); - } - Ok(()) -} - -async fn run_tesseract( - executable: &Path, - source: &Path, - languages: &[String], - page_segmentation_mode: Option, - timeout: std::time::Duration, -) -> UseResult> { - let mut command = Command::new(executable); - command - .arg(source) - .arg("stdout") - .arg("-l") - .arg(languages.join("+")) - .stdout(Stdio::piped()) - .stderr(Stdio::piped()) - .kill_on_drop(true); - if let Some(mode) = page_segmentation_mode { - command.arg("--psm").arg(mode.to_string()); - } - command.arg("tsv"); - - let output = tokio::time::timeout(timeout, command.output()) - .await - .map_err(|_| { - UseError::new( - "use.ocr.provider_timeout", - format!( - "Tesseract exceeded the {} ms OCR timeout.", - timeout.as_millis() - ), - ) - })? - .map_err(|error| { - UseError::new( - "use.ocr.provider_failed", - format!( - "Failed to launch Tesseract executable '{}': {error}", - executable.display() - ), - ) - })?; - if output.stdout.len() > MAX_PROVIDER_OUTPUT_BYTES - || output.stderr.len() > MAX_PROVIDER_OUTPUT_BYTES - { - return Err(UseError::new( - "use.ocr.output_too_large", - "Tesseract output exceeded 8 MiB.", - )); - } - if !output.status.success() { - let stderr = String::from_utf8_lossy(&output.stderr); - return Err(UseError::new( - "use.ocr.provider_failed", - format!( - "Tesseract exited with {}: {}", - output.status, - bounded_text(&stderr, 2048) - ), - )); - } - Ok(output.stdout) -} - -#[derive(Default)] -struct LineAccumulator { - page: u32, - words: Vec, - confidence_sum: f32, - confidence_count: usize, - left: u32, - top: u32, - right: u32, - bottom: u32, -} - -fn parse_tesseract_tsv(output: &[u8]) -> UseResult<(String, Vec)> { - let output = std::str::from_utf8(output).map_err(|error| { - UseError::new( - "use.ocr.provider_output_invalid", - format!("Tesseract TSV output was not UTF-8: {error}"), - ) - })?; - let mut lines = BTreeMap::<(u32, u32, u32, u32), LineAccumulator>::new(); - for (index, row) in output.lines().enumerate() { - if index == 0 && row.starts_with("level\t") { - continue; - } - if row.trim().is_empty() { - continue; - } - let columns = row.splitn(12, '\t').collect::>(); - if columns.len() != 12 { - return Err(UseError::new( - "use.ocr.provider_output_invalid", - format!( - "Tesseract TSV row {} did not contain 12 columns.", - index + 1 - ), - )); - } - let level = parse_u32(columns[0], index)?; - if level != 5 || columns[11].trim().is_empty() { - continue; - } - let page = parse_u32(columns[1], index)?; - let block = parse_u32(columns[2], index)?; - let paragraph = parse_u32(columns[3], index)?; - let line = parse_u32(columns[4], index)?; - let left = parse_u32(columns[6], index)?; - let top = parse_u32(columns[7], index)?; - let width = parse_u32(columns[8], index)?; - let height = parse_u32(columns[9], index)?; - let confidence = columns[10] - .parse::() - .ok() - .filter(|value| *value >= 0.0); - let entry = lines - .entry((page, block, paragraph, line)) - .or_insert_with(|| LineAccumulator { - page, - left, - top, - right: left.saturating_add(width), - bottom: top.saturating_add(height), - ..LineAccumulator::default() - }); - entry.words.push(columns[11].trim().to_string()); - if let Some(confidence) = confidence { - entry.confidence_sum += confidence; - entry.confidence_count += 1; - } - entry.left = entry.left.min(left); - entry.top = entry.top.min(top); - entry.right = entry.right.max(left.saturating_add(width)); - entry.bottom = entry.bottom.max(top.saturating_add(height)); - } - let blocks = lines - .into_values() - .filter_map(|line| { - let text = line.words.join(" "); - (!text.is_empty()).then(|| OcrBlock { - page: line.page, - text, - confidence: (line.confidence_count > 0) - .then(|| line.confidence_sum / line.confidence_count as f32), - bounding_box: Some(OcrBoundingBox { - x: line.left, - y: line.top, - width: line.right.saturating_sub(line.left), - height: line.bottom.saturating_sub(line.top), - }), - }) - }) - .collect::>(); - let text = blocks - .iter() - .map(|block| block.text.as_str()) - .collect::>() - .join("\n"); - Ok((text, blocks)) -} - -fn parse_u32(value: &str, row: usize) -> UseResult { - value.parse::().map_err(|_| { - UseError::new( - "use.ocr.provider_output_invalid", - format!( - "Tesseract TSV row {} contained an invalid integer.", - row + 1 - ), - ) - }) -} - -fn vision_content_text(content: &serde_json::Value) -> UseResult { - if let Some(text) = content.as_str() { - return Ok(text.to_string()); - } - let Some(parts) = content.as_array() else { - return Err(UseError::new( - "use.ocr.vision_response_invalid", - "Vision OCR message content was neither text nor a text-part array.", - )); - }; - let text = parts - .iter() - .filter_map(|part| { - part.get("text") - .and_then(serde_json::Value::as_str) - .or_else(|| part.as_str()) - }) - .collect::>() - .join(""); - if text.is_empty() { - return Err(UseError::new( - "use.ocr.vision_response_invalid", - "Vision OCR message content did not contain text.", - )); - } - Ok(text) -} - fn detect_image_type(bytes: &[u8]) -> Option<&'static str> { if bytes.starts_with(b"\x89PNG\r\n\x1a\n") { Some("image/png") @@ -562,46 +283,10 @@ fn detect_image_type(bytes: &[u8]) -> Option<&'static str> { } } -fn bounded_text(value: &str, max: usize) -> String { - let mut text = value.chars().take(max).collect::(); - if value.chars().count() > max { - text.push('…'); - } - text -} - -fn redacted_endpoint(endpoint: &url::Url) -> String { - let mut redacted = endpoint.clone(); - redacted.set_query(None); - redacted.set_fragment(None); - redacted.to_string() -} - #[cfg(test)] mod tests { use super::*; - #[cfg(unix)] - use std::os::unix::fs::PermissionsExt; - - #[test] - fn parses_tesseract_words_into_ordered_lines() { - let tsv = b"level\tpage_num\tblock_num\tpar_num\tline_num\tword_num\tleft\ttop\twidth\theight\tconf\ttext\n5\t1\t1\t1\t1\t1\t10\t20\t30\t10\t95.0\tHello\n5\t1\t1\t1\t1\t2\t45\t20\t35\t10\t85.0\tworld\n5\t1\t1\t1\t2\t1\t10\t40\t20\t10\t90.0\tNext\n"; - let (text, blocks) = parse_tesseract_tsv(tsv).unwrap(); - assert_eq!(text, "Hello world\nNext"); - assert_eq!(blocks.len(), 2); - assert_eq!(blocks[0].confidence, Some(90.0)); - assert_eq!( - blocks[0].bounding_box, - Some(OcrBoundingBox { - x: 10, - y: 20, - width: 70, - height: 10, - }) - ); - } - #[test] fn detects_supported_image_signatures() { assert_eq!( @@ -612,38 +297,11 @@ mod tests { assert_eq!(detect_image_type(b"not an image"), None); } - #[cfg(unix)] - #[tokio::test] - async fn local_provider_extracts_a_real_bounded_source_through_its_process_boundary() { - let temp = tempfile::tempdir().unwrap(); - let executable = temp.path().join("tesseract-fixture"); - std::fs::write( - &executable, - "#!/bin/sh\nprintf 'level\\tpage_num\\tblock_num\\tpar_num\\tline_num\\tword_num\\tleft\\ttop\\twidth\\theight\\tconf\\ttext\\n5\\t1\\t1\\t1\\t1\\t1\\t2\\t3\\t20\\t8\\t98.0\\tA3S\\n5\\t1\\t1\\t1\\t1\\t2\\t24\\t3\\t30\\t8\\t96.0\\tUse\\n'\n", - ) - .unwrap(); - let mut permissions = std::fs::metadata(&executable).unwrap().permissions(); - permissions.set_mode(0o755); - std::fs::set_permissions(&executable, permissions).unwrap(); - let image = temp.path().join("scan.png"); - std::fs::write(&image, b"\x89PNG\r\n\x1a\nfixture").unwrap(); - - let result = OcrClient::with_tesseract(executable) - .unwrap() - .extract(OcrRequest { - path: image, - languages: vec!["eng".to_string()], - page_segmentation_mode: Some(6), - provider: Some(OcrProviderKind::Tesseract), - prompt: None, - }) - .await - .unwrap(); - - assert_eq!(result.provider, OcrProviderKind::Tesseract); - assert_eq!(result.text, "A3S Use"); - assert_eq!(result.blocks.len(), 1); - assert_eq!(result.source.media_type, "image/png"); - assert_eq!(result.source.sha256.len(), 64); + #[test] + fn diagnostic_never_discloses_an_off_device_provider() { + let diagnostic = OcrClient::from_env().unwrap().diagnostic(); + assert_eq!(diagnostic.provider, Some(OcrProviderKind::PpOcrV6)); + assert!(!diagnostic.sends_source_off_device); + assert_eq!(diagnostic.engine.as_deref(), Some(ENGINE_NAME)); } } diff --git a/crates/ocr/src/config.rs b/crates/ocr/src/config.rs new file mode 100644 index 00000000..9395ed07 --- /dev/null +++ b/crates/ocr/src/config.rs @@ -0,0 +1,261 @@ +use std::collections::BTreeMap; +use std::path::Path; + +use a3s_use_core::{UseError, UseResult}; +use serde::Deserialize; +use serde_yaml::Value; + +pub(crate) const MODEL_FAMILY: &str = "PP-OCRv6_small"; +pub(crate) const DETECTION_MODEL: &str = "PP-OCRv6_small_det"; +pub(crate) const RECOGNITION_MODEL: &str = "PP-OCRv6_small_rec"; + +#[derive(Debug, Clone)] +pub(crate) struct DetectionConfig { + pub(crate) scale: f32, + pub(crate) mean: [f32; 3], + pub(crate) std: [f32; 3], + pub(crate) threshold: f32, + pub(crate) box_threshold: f32, + pub(crate) max_candidates: usize, + pub(crate) unclip_ratio: f32, +} + +#[derive(Debug, Clone)] +pub(crate) struct RecognitionConfig { + pub(crate) channels: usize, + pub(crate) height: usize, + pub(crate) default_width: usize, + pub(crate) characters: Vec, +} + +#[derive(Debug, Deserialize)] +struct RawConfig { + #[serde(rename = "Global")] + global: RawGlobal, + #[serde(rename = "PreProcess")] + pre_process: RawPreProcess, + #[serde(rename = "PostProcess")] + post_process: RawPostProcess, +} + +#[derive(Debug, Deserialize)] +struct RawGlobal { + model_name: String, +} + +#[derive(Debug, Deserialize)] +struct RawPreProcess { + transform_ops: Vec>, +} + +#[derive(Debug, Deserialize)] +struct RawPostProcess { + #[serde(default)] + name: String, + #[serde(default = "default_threshold")] + thresh: f32, + #[serde(default = "default_box_threshold")] + box_thresh: f32, + #[serde(default = "default_max_candidates")] + max_candidates: usize, + #[serde(default = "default_unclip_ratio")] + unclip_ratio: f32, + #[serde(default)] + character_dict: Vec, +} + +pub(crate) fn load_detection(path: &Path) -> UseResult { + let raw = load(path)?; + if raw.global.model_name != DETECTION_MODEL { + return Err(config_error(format!( + "Expected detection model '{DETECTION_MODEL}', found '{}'.", + raw.global.model_name + ))); + } + if raw.post_process.name != "DBPostProcess" { + return Err(config_error(format!( + "Expected DBPostProcess, found '{}'.", + raw.post_process.name + ))); + } + let normalize = transform(&raw.pre_process.transform_ops, "NormalizeImage") + .ok_or_else(|| config_error("Detection config has no NormalizeImage transform."))?; + let scale = normalize + .get("scale") + .and_then(parse_scale) + .unwrap_or(1.0 / 255.0); + let mean = float_triplet(normalize.get("mean"), [0.485, 0.456, 0.406])?; + let std = float_triplet(normalize.get("std"), [0.229, 0.224, 0.225])?; + if std.iter().any(|value| *value <= 0.0) { + return Err(config_error( + "Detection normalization standard deviations must be positive.", + )); + } + Ok(DetectionConfig { + scale, + mean, + std, + threshold: raw.post_process.thresh, + box_threshold: raw.post_process.box_thresh, + max_candidates: raw.post_process.max_candidates.min(10_000), + unclip_ratio: raw.post_process.unclip_ratio, + }) +} + +pub(crate) fn load_recognition(path: &Path) -> UseResult { + let raw = load(path)?; + if raw.global.model_name != RECOGNITION_MODEL { + return Err(config_error(format!( + "Expected recognition model '{RECOGNITION_MODEL}', found '{}'.", + raw.global.model_name + ))); + } + if raw.post_process.name != "CTCLabelDecode" { + return Err(config_error(format!( + "Expected CTCLabelDecode, found '{}'.", + raw.post_process.name + ))); + } + if raw.post_process.character_dict.is_empty() || raw.post_process.character_dict.len() > 100_000 + { + return Err(config_error( + "Recognition character dictionary is empty or unreasonably large.", + )); + } + let resize = transform(&raw.pre_process.transform_ops, "RecResizeImg") + .ok_or_else(|| config_error("Recognition config has no RecResizeImg transform."))?; + let shape = resize + .get("image_shape") + .and_then(Value::as_sequence) + .ok_or_else(|| config_error("RecResizeImg.image_shape must be an integer triplet."))?; + if shape.len() != 3 { + return Err(config_error( + "RecResizeImg.image_shape must contain channels, height, and width.", + )); + } + let dimensions = shape + .iter() + .map(|value| value.as_u64().and_then(|value| usize::try_from(value).ok())) + .collect::>>() + .ok_or_else(|| config_error("RecResizeImg.image_shape contains an invalid dimension."))?; + if dimensions[0] != 3 + || !(16..=256).contains(&dimensions[1]) + || !(32..=4096).contains(&dimensions[2]) + { + return Err(config_error(format!( + "Unsupported PP-OCRv6 recognition input shape {:?}.", + dimensions + ))); + } + Ok(RecognitionConfig { + channels: dimensions[0], + height: dimensions[1], + default_width: dimensions[2], + characters: raw.post_process.character_dict, + }) +} + +fn load(path: &Path) -> UseResult { + let metadata = std::fs::metadata(path).map_err(|error| { + config_error(format!( + "Failed to inspect PP-OCRv6 config '{}': {error}", + path.display() + )) + })?; + if !metadata.is_file() || metadata.len() == 0 || metadata.len() > 2 * 1024 * 1024 { + return Err(config_error(format!( + "PP-OCRv6 config '{}' must be a non-empty regular file no larger than 2 MiB.", + path.display() + ))); + } + let text = std::fs::read_to_string(path).map_err(|error| { + config_error(format!( + "Failed to read PP-OCRv6 config '{}': {error}", + path.display() + )) + })?; + serde_yaml::from_str(&text).map_err(|error| { + config_error(format!( + "Failed to parse PP-OCRv6 config '{}': {error}", + path.display() + )) + }) +} + +fn transform<'a>( + transforms: &'a [BTreeMap], + name: &str, +) -> Option<&'a serde_yaml::Mapping> { + transforms + .iter() + .find_map(|transform| transform.get(name)) + .and_then(Value::as_mapping) +} + +fn float_triplet(value: Option<&Value>, default: [f32; 3]) -> UseResult<[f32; 3]> { + let Some(values) = value.and_then(Value::as_sequence) else { + return Ok(default); + }; + if values.len() != 3 { + return Err(config_error( + "Detection normalization mean and std must contain three values.", + )); + } + let mut output = [0.0_f32; 3]; + for (index, value) in values.iter().enumerate() { + output[index] = yaml_f32(value) + .ok_or_else(|| config_error("Detection normalization contains a non-number."))?; + } + Ok(output) +} + +fn parse_scale(value: &Value) -> Option { + if let Some(value) = yaml_f32(value) { + return Some(value); + } + let value = value.as_str()?.trim(); + if let Some((numerator, denominator)) = value.split_once('/') { + let numerator = numerator.trim_matches('.').parse::().ok()?; + let denominator = denominator.trim_matches('.').parse::().ok()?; + return (denominator != 0.0).then_some(numerator / denominator); + } + value.parse().ok() +} + +fn yaml_f32(value: &Value) -> Option { + value + .as_f64() + .map(|value| value as f32) + .filter(|value| value.is_finite()) +} + +fn default_threshold() -> f32 { + 0.3 +} + +fn default_box_threshold() -> f32 { + 0.6 +} + +fn default_max_candidates() -> usize { + 1_000 +} + +fn default_unclip_ratio() -> f32 { + 1.5 +} + +fn config_error(message: impl Into) -> UseError { + UseError::new("use.ocr.model_config_invalid", message) +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn parses_fractional_detection_scale() { + let value = Value::String("1./255.".to_string()); + assert!((parse_scale(&value).unwrap() - 1.0 / 255.0).abs() < f32::EPSILON); + } +} diff --git a/crates/ocr/src/engine.rs b/crates/ocr/src/engine.rs new file mode 100644 index 00000000..4c72f1fa --- /dev/null +++ b/crates/ocr/src/engine.rs @@ -0,0 +1,261 @@ +use std::path::Path; + +use a3s_use_core::{UseError, UseResult}; +use image::{imageops, ImageBuffer, Rgb, RgbImage}; +use imageproc::geometric_transformations::{warp_into, Interpolation, Projection}; +use imageproc::point::Point; +use ort::session::builder::GraphOptimizationLevel; +use ort::session::Session; +use ort::value::TensorRef; + +use crate::assets::ModelAssets; +use crate::config::{load_detection, load_recognition, DetectionConfig, RecognitionConfig}; +use crate::postprocess::{decode_ctc, detection_boxes, Detection}; +use crate::preprocess::{detection_input, recognition_input}; + +const RECOGNITION_BATCH_SIZE: usize = 8; +const MAX_CROP_PIXELS: u64 = 64 * 1024 * 1024; + +#[derive(Debug, Clone)] +pub(crate) struct EngineBlock { + pub(crate) polygon: [Point; 4], + pub(crate) detection_confidence: f32, + pub(crate) text: String, + pub(crate) confidence: f32, +} + +pub(crate) struct PpOcrV6Engine { + detection: Session, + recognition: Session, + detection_config: DetectionConfig, + recognition_config: RecognitionConfig, +} + +impl PpOcrV6Engine { + pub(crate) fn load(assets: &ModelAssets) -> UseResult { + let detection_config = load_detection(&assets.detection_config)?; + let recognition_config = load_recognition(&assets.recognition_config)?; + let detection = load_session(&assets.detection_model, "detection")?; + let recognition = load_session(&assets.recognition_model, "recognition")?; + Ok(Self { + detection, + recognition, + detection_config, + recognition_config, + }) + } + + pub(crate) fn extract(&mut self, image: &RgbImage) -> UseResult> { + let input = detection_input(image, &self.detection_config)?; + let (shape, output) = + run_session(&mut self.detection, &input.data, input.shape, "detection")?; + let detections = detection_boxes( + &output, + &shape, + input.original_width, + input.original_height, + &self.detection_config, + )?; + if detections.is_empty() { + return Ok(Vec::new()); + } + + let crops = detections + .iter() + .map(|detection| perspective_crop(image, detection)) + .collect::>>()?; + let mut blocks = Vec::with_capacity(detections.len()); + for (detection_batch, crop_batch) in detections + .chunks(RECOGNITION_BATCH_SIZE) + .zip(crops.chunks(RECOGNITION_BATCH_SIZE)) + { + let input = recognition_input(crop_batch, &self.recognition_config)?; + let (shape, output) = run_session( + &mut self.recognition, + &input.data, + input.shape, + "recognition", + )?; + if shape.len() != 3 || shape[0] != detection_batch.len() { + return Err(engine_error( + "use.ocr.provider_output_invalid", + format!( + "PP-OCRv6 recognition output shape must be [N, T, C] for N={}, found {shape:?}.", + detection_batch.len() + ), + )); + } + let item_len = shape[1].checked_mul(shape[2]).ok_or_else(|| { + engine_error( + "use.ocr.provider_output_invalid", + "PP-OCRv6 recognition output dimensions overflowed.", + ) + })?; + if output.len() != detection_batch.len().saturating_mul(item_len) { + return Err(engine_error( + "use.ocr.provider_output_invalid", + "PP-OCRv6 recognition output length does not match its batch shape.", + )); + } + for (index, detection) in detection_batch.iter().enumerate() { + let start = index * item_len; + let recognition = decode_ctc( + &output[start..start + item_len], + &[1, shape[1], shape[2]], + &self.recognition_config, + )?; + blocks.push(EngineBlock { + polygon: detection.polygon, + detection_confidence: detection.confidence, + text: recognition.text, + confidence: recognition.confidence, + }); + } + } + Ok(blocks) + } +} + +fn load_session(path: &Path, role: &str) -> UseResult { + let session = Session::builder() + .map_err(|error| runtime_error(role, "create an ONNX Runtime session", error))? + .with_optimization_level(GraphOptimizationLevel::Level3) + .map_err(|error| runtime_error(role, "configure graph optimization", error))? + .commit_from_file(path) + .map_err(|error| runtime_error(role, "load the ONNX model", error))?; + if session.inputs.len() != 1 || session.outputs.len() != 1 { + return Err(engine_error( + "use.ocr.model_invalid", + format!( + "PP-OCRv6 {role} model must expose exactly one input and one output; found {} inputs and {} outputs.", + session.inputs.len(), + session.outputs.len() + ), + )); + } + Ok(session) +} + +fn run_session( + session: &mut Session, + data: &[f32], + shape: [usize; 4], + role: &str, +) -> UseResult<(Vec, Vec)> { + let expected = shape + .iter() + .try_fold(1_usize, |total, dimension| total.checked_mul(*dimension)); + if expected != Some(data.len()) { + return Err(engine_error( + "use.ocr.provider_input_invalid", + format!("PP-OCRv6 {role} tensor length does not match its shape."), + )); + } + let input = TensorRef::from_array_view((shape, data)) + .map_err(|error| runtime_error(role, "create an ONNX Runtime input tensor", error))?; + let outputs = session + .run(ort::inputs![input]) + .map_err(|error| runtime_error(role, "run ONNX inference", error))?; + if outputs.len() != 1 { + return Err(engine_error( + "use.ocr.provider_output_invalid", + format!( + "PP-OCRv6 {role} inference returned {} outputs instead of one.", + outputs.len() + ), + )); + } + let output = outputs.values().next().ok_or_else(|| { + engine_error( + "use.ocr.provider_output_invalid", + format!("PP-OCRv6 {role} inference returned no output tensor."), + ) + })?; + let (output_shape, output_data) = output + .try_extract_tensor::() + .map_err(|error| runtime_error(role, "read the ONNX output tensor", error))?; + let output_shape = output_shape + .iter() + .map(|dimension| { + usize::try_from(*dimension).map_err(|_| { + engine_error( + "use.ocr.provider_output_invalid", + format!("PP-OCRv6 {role} output contains an invalid dimension {dimension}."), + ) + }) + }) + .collect::>>()?; + if output_data.iter().any(|value| !value.is_finite()) { + return Err(engine_error( + "use.ocr.provider_output_invalid", + format!("PP-OCRv6 {role} output contains a non-finite value."), + )); + } + Ok((output_shape, output_data.to_vec())) +} + +fn perspective_crop(image: &RgbImage, detection: &Detection) -> UseResult { + let polygon = detection.polygon; + let width = distance(polygon[0], polygon[1]) + .max(distance(polygon[2], polygon[3])) + .round() + .max(1.0) as u32; + let height = distance(polygon[0], polygon[3]) + .max(distance(polygon[1], polygon[2])) + .round() + .max(1.0) as u32; + let pixels = u64::from(width) + .checked_mul(u64::from(height)) + .ok_or_else(|| crop_error("PP-OCRv6 text crop dimensions overflowed."))?; + if pixels > MAX_CROP_PIXELS { + return Err(crop_error( + "PP-OCRv6 text crop exceeds the 64 megapixel safety limit.", + )); + } + + let source = polygon.map(|point| (point.x, point.y)); + let destination = [ + (0.0, 0.0), + (width.saturating_sub(1) as f32, 0.0), + ( + width.saturating_sub(1) as f32, + height.saturating_sub(1) as f32, + ), + (0.0, height.saturating_sub(1) as f32), + ]; + let projection = Projection::from_control_points(source, destination).ok_or_else(|| { + crop_error("PP-OCRv6 detected a degenerate text polygon that cannot be rectified.") + })?; + let mut crop = ImageBuffer::new(width, height); + warp_into( + image, + &projection, + Interpolation::Bicubic, + Rgb([255, 255, 255]), + &mut crop, + ); + if f64::from(height) / f64::from(width) >= 1.5 { + Ok(imageops::rotate270(&crop)) + } else { + Ok(crop) + } +} + +fn distance(left: Point, right: Point) -> f32 { + (left.x - right.x).hypot(left.y - right.y) +} + +fn runtime_error(role: &str, action: &str, error: impl std::fmt::Display) -> UseError { + engine_error( + "use.ocr.runtime_failed", + format!("Failed to {action} for PP-OCRv6 {role}: {error}"), + ) +} + +fn crop_error(message: impl Into) -> UseError { + engine_error("use.ocr.crop_invalid", message) +} + +fn engine_error(code: &str, message: impl Into) -> UseError { + UseError::new(code, message) +} diff --git a/crates/ocr/src/install.rs b/crates/ocr/src/install.rs new file mode 100644 index 00000000..91031170 --- /dev/null +++ b/crates/ocr/src/install.rs @@ -0,0 +1,671 @@ +use std::fs::OpenOptions; +use std::io::{Read, Write}; +use std::path::{Component, Path, PathBuf}; +use std::sync::atomic::{AtomicU64, Ordering}; + +use a3s_use_core::{UseError, UseResult}; +use fs2::FileExt; +use serde::{Deserialize, Serialize}; +use sha2::{Digest, Sha256}; +use tokio::io::AsyncWriteExt; + +use crate::assets::{ + managed_model_dir, managed_root, ocr_status, validate_assets, OcrInstallSource, + OcrRuntimeStatus, RECEIPT_FILE, +}; +use crate::config::MODEL_FAMILY; + +const INSTALL_LOCK: &str = ".install.lock"; +const STAGE_PREFIX: &str = ".stage-"; +const BACKUP_PREFIX: &str = ".backup-"; +const DOWNLOAD_HOST: &str = "paddle-model-ecology.bj.bcebos.com"; +const MAX_ARCHIVE_BYTES: u64 = 256 * 1024 * 1024; + +const DETECTION_ARCHIVE: PinnedArchive = PinnedArchive { + role: "det", + directory: "PP-OCRv6_small_det_onnx_infer", + url: "https://paddle-model-ecology.bj.bcebos.com/paddlex/official_inference_model/paddle3.0.0/PP-OCRv6_small_det_onnx_infer.tar", + bytes: 9_891_840, + sha256: "d218f6fbf0f1c23d2161bd6ac7f5eaa6104fa89955c09290497e31008e2618e4", +}; +const RECOGNITION_ARCHIVE: PinnedArchive = PinnedArchive { + role: "rec", + directory: "PP-OCRv6_small_rec_onnx_infer", + url: "https://paddle-model-ecology.bj.bcebos.com/paddlex/official_inference_model/paddle3.0.0/PP-OCRv6_small_rec_onnx_infer.tar", + bytes: 21_319_680, + sha256: "d267ab077a44a0eedb1ea8f8c542d263f211de8e9d7a029bf9fcfff7e5a88fb1", +}; + +#[derive(Debug, Clone, Copy)] +struct PinnedArchive { + role: &'static str, + directory: &'static str, + url: &'static str, + bytes: u64, + sha256: &'static str, +} + +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +#[serde(rename_all = "camelCase")] +struct InstallReceipt { + schema_version: u32, + provider: String, + model: String, + detection_url: String, + detection_sha256: String, + recognition_url: String, + recognition_sha256: String, +} + +struct InstallLock { + _file: std::fs::File, +} + +struct Downloaded { + bytes: u64, + sha256: String, +} + +pub async fn install_ppocr_v6(force: bool) -> UseResult { + let current = ocr_status(); + if !force && current.available { + return Ok(current); + } + + let root = managed_root()?; + let _lock = acquire_lock(&root).await?; + cleanup_stale(&root).await?; + + let current = ocr_status(); + if !force && current.available { + return Ok(current); + } + + let stage = create_stage(&root).await?; + let install_result = install_into(&stage).await; + if install_result.is_err() { + let _ = tokio::fs::remove_dir_all(&stage).await; + } + install_result?; + + let target = managed_model_dir()?; + activate(&stage, &target).await?; + validate_assets(&target, OcrInstallSource::Managed)?; + + let status = ocr_status(); + if status.available { + Ok(status) + } else { + Err(ocr_error( + "use.ocr.install_failed", + "PP-OCRv6 installation completed without a usable model bundle.", + )) + } +} + +pub async fn repair_ppocr_v6() -> UseResult { + let status = ocr_status(); + if status.available { + Ok(status) + } else { + install_ppocr_v6(true).await + } +} + +pub async fn uninstall_managed_ppocr_v6() -> UseResult { + let root = managed_root()?; + let _lock = acquire_lock(&root).await?; + let target = managed_model_dir()?; + if !owned_install(&target) { + return Ok(false); + } + tokio::fs::remove_dir_all(&target).await.map_err(|error| { + ocr_error( + "use.ocr.uninstall_failed", + format!( + "Failed to remove managed PP-OCRv6 bundle '{}': {error}", + target.display() + ), + ) + })?; + Ok(true) +} + +async fn install_into(stage: &Path) -> UseResult<()> { + let client = download_client()?; + for archive in [DETECTION_ARCHIVE, RECOGNITION_ARCHIVE] { + let archive_path = stage.join(format!("{}.tar", archive.role)); + let downloaded = download(&client, archive.url, &archive_path).await?; + if downloaded.bytes != archive.bytes || downloaded.sha256 != archive.sha256 { + return Err(ocr_error( + "use.ocr.integrity_mismatch", + format!( + "{} archive integrity mismatch: expected {} bytes and {}, got {} bytes and {}.", + archive.directory, + archive.bytes, + archive.sha256, + downloaded.bytes, + downloaded.sha256 + ), + )); + } + let archive_path_for_task = archive_path.clone(); + let destination = stage.join(archive.role); + tokio::task::spawn_blocking(move || { + extract_archive(&archive_path_for_task, &destination, archive) + }) + .await + .map_err(|error| { + ocr_error( + "use.ocr.install_failed", + format!("PP-OCRv6 archive extraction task failed: {error}"), + ) + })??; + tokio::fs::remove_file(&archive_path) + .await + .map_err(|error| { + ocr_error( + "use.ocr.install_failed", + format!( + "Failed to remove staged archive '{}': {error}", + archive_path.display() + ), + ) + })?; + } + write_receipt(stage).await?; + validate_assets(stage, OcrInstallSource::Managed)?; + Ok(()) +} + +fn download_client() -> UseResult { + let redirects = reqwest::redirect::Policy::custom(|attempt| { + let approved = attempt.previous().len() < 5 + && attempt.url().scheme() == "https" + && attempt.url().host_str() == Some(DOWNLOAD_HOST); + if approved { + attempt.follow() + } else { + attempt.error("PP-OCRv6 download redirected to an unapproved host") + } + }); + reqwest::Client::builder() + .user_agent(concat!("a3s-use-ocr/", env!("CARGO_PKG_VERSION"))) + .redirect(redirects) + .timeout(std::time::Duration::from_secs(300)) + .build() + .map_err(|error| { + ocr_error( + "use.ocr.download_failed", + format!("Failed to create PP-OCRv6 download client: {error}"), + ) + }) +} + +async fn download( + client: &reqwest::Client, + value: &str, + destination: &Path, +) -> UseResult { + let url = reqwest::Url::parse(value).map_err(|error| { + ocr_error( + "use.ocr.download_source_invalid", + format!("Invalid PP-OCRv6 download URL: {error}"), + ) + })?; + if url.scheme() != "https" || url.host_str() != Some(DOWNLOAD_HOST) { + return Err(ocr_error( + "use.ocr.download_source_invalid", + "PP-OCRv6 download source is not the pinned official HTTPS host.", + )); + } + let mut response = client + .get(url) + .send() + .await + .map_err(|error| { + ocr_error( + "use.ocr.download_failed", + format!("Failed to download PP-OCRv6: {error}"), + ) + })? + .error_for_status() + .map_err(|error| { + ocr_error( + "use.ocr.download_failed", + format!("PP-OCRv6 download failed: {error}"), + ) + })?; + if response + .content_length() + .is_some_and(|length| length > MAX_ARCHIVE_BYTES) + { + return Err(ocr_error( + "use.ocr.download_too_large", + "PP-OCRv6 archive exceeds the 256 MiB limit.", + )); + } + let mut file = tokio::fs::OpenOptions::new() + .create_new(true) + .write(true) + .open(destination) + .await + .map_err(|error| { + ocr_error( + "use.ocr.install_failed", + format!( + "Failed to create PP-OCRv6 download '{}': {error}", + destination.display() + ), + ) + })?; + let mut hasher = Sha256::new(); + let mut total = 0_u64; + while let Some(chunk) = response.chunk().await.map_err(|error| { + ocr_error( + "use.ocr.download_failed", + format!("Failed to read PP-OCRv6 download: {error}"), + ) + })? { + total = total + .checked_add(chunk.len() as u64) + .ok_or_else(|| ocr_error("use.ocr.download_too_large", "Download size overflowed."))?; + if total > MAX_ARCHIVE_BYTES { + return Err(ocr_error( + "use.ocr.download_too_large", + "PP-OCRv6 archive exceeds the 256 MiB limit.", + )); + } + hasher.update(&chunk); + file.write_all(&chunk).await.map_err(|error| { + ocr_error( + "use.ocr.install_failed", + format!( + "Failed to write PP-OCRv6 download '{}': {error}", + destination.display() + ), + ) + })?; + } + file.flush().await.map_err(|error| { + ocr_error( + "use.ocr.install_failed", + format!( + "Failed to flush PP-OCRv6 download '{}': {error}", + destination.display() + ), + ) + })?; + file.sync_all().await.map_err(|error| { + ocr_error( + "use.ocr.install_failed", + format!( + "Failed to sync PP-OCRv6 download '{}': {error}", + destination.display() + ), + ) + })?; + Ok(Downloaded { + bytes: total, + sha256: format!("{:x}", hasher.finalize()), + }) +} + +fn extract_archive(archive_path: &Path, destination: &Path, spec: PinnedArchive) -> UseResult<()> { + std::fs::create_dir(destination).map_err(|error| { + ocr_error( + "use.ocr.install_failed", + format!( + "Failed to create PP-OCRv6 model directory '{}': {error}", + destination.display() + ), + ) + })?; + let file = std::fs::File::open(archive_path).map_err(|error| { + ocr_error( + "use.ocr.install_failed", + format!( + "Failed to open PP-OCRv6 archive '{}': {error}", + archive_path.display() + ), + ) + })?; + let mut archive = tar::Archive::new(file); + let mut extracted = [false; 2]; + for entry in archive.entries().map_err(archive_error)? { + let entry = entry.map_err(archive_error)?; + let path = entry.path().map_err(archive_error)?; + let components = path.components().collect::>(); + if components.len() == 1 + && matches!(components[0], Component::Normal(value) if value == spec.directory) + && entry.header().entry_type().is_dir() + { + continue; + } + if components.len() != 2 + || !matches!(components[0], Component::Normal(value) if value == spec.directory) + || !entry.header().entry_type().is_file() + { + return Err(ocr_error( + "use.ocr.archive_invalid", + format!( + "PP-OCRv6 archive contains an unexpected entry '{}'.", + path.display() + ), + )); + } + let name = match components[1] { + Component::Normal(name) if name == "inference.onnx" => { + extracted[0] = true; + "inference.onnx" + } + Component::Normal(name) if name == "inference.yml" => { + extracted[1] = true; + "inference.yml" + } + _ => { + return Err(ocr_error( + "use.ocr.archive_invalid", + format!( + "PP-OCRv6 archive contains an unexpected entry '{}'.", + path.display() + ), + )) + } + }; + let max = if name.ends_with(".onnx") { + 256 * 1024 * 1024 + } else { + 2 * 1024 * 1024 + }; + if entry.size() == 0 || entry.size() > max { + return Err(ocr_error( + "use.ocr.archive_invalid", + format!("PP-OCRv6 archive entry '{name}' has an invalid size."), + )); + } + let expected_size = entry.size(); + let output_path = destination.join(name); + let mut output = OpenOptions::new() + .create_new(true) + .write(true) + .open(&output_path) + .map_err(|error| { + ocr_error( + "use.ocr.install_failed", + format!( + "Failed to create PP-OCRv6 asset '{}': {error}", + output_path.display() + ), + ) + })?; + let copied = std::io::copy(&mut entry.take(max + 1), &mut output).map_err(archive_error)?; + if copied != expected_size { + return Err(ocr_error( + "use.ocr.archive_invalid", + format!("PP-OCRv6 archive entry '{name}' was truncated."), + )); + } + output.flush().map_err(archive_error)?; + output.sync_all().map_err(archive_error)?; + } + if !extracted.into_iter().all(|present| present) { + return Err(ocr_error( + "use.ocr.archive_invalid", + "PP-OCRv6 archive is missing inference.onnx or inference.yml.", + )); + } + Ok(()) +} + +async fn write_receipt(stage: &Path) -> UseResult<()> { + let receipt = InstallReceipt { + schema_version: 1, + provider: "pp-ocr-v6".to_string(), + model: MODEL_FAMILY.to_string(), + detection_url: DETECTION_ARCHIVE.url.to_string(), + detection_sha256: DETECTION_ARCHIVE.sha256.to_string(), + recognition_url: RECOGNITION_ARCHIVE.url.to_string(), + recognition_sha256: RECOGNITION_ARCHIVE.sha256.to_string(), + }; + let bytes = serde_json::to_vec_pretty(&receipt).map_err(|error| { + ocr_error( + "use.ocr.install_failed", + format!("Failed to encode PP-OCRv6 install receipt: {error}"), + ) + })?; + let path = stage.join(RECEIPT_FILE); + let mut file = tokio::fs::OpenOptions::new() + .create_new(true) + .write(true) + .open(&path) + .await + .map_err(|error| { + ocr_error( + "use.ocr.install_failed", + format!( + "Failed to create PP-OCRv6 receipt '{}': {error}", + path.display() + ), + ) + })?; + file.write_all(&bytes).await.map_err(|error| { + ocr_error( + "use.ocr.install_failed", + format!( + "Failed to write PP-OCRv6 receipt '{}': {error}", + path.display() + ), + ) + })?; + file.sync_all().await.map_err(|error| { + ocr_error( + "use.ocr.install_failed", + format!( + "Failed to sync PP-OCRv6 receipt '{}': {error}", + path.display() + ), + ) + }) +} + +async fn acquire_lock(root: &Path) -> UseResult { + tokio::fs::create_dir_all(root).await.map_err(|error| { + ocr_error( + "use.ocr.install_failed", + format!( + "Failed to create OCR data root '{}': {error}", + root.display() + ), + ) + })?; + let path = root.join(INSTALL_LOCK); + tokio::task::spawn_blocking(move || { + let file = OpenOptions::new() + .create(true) + .read(true) + .write(true) + .truncate(false) + .open(&path) + .map_err(|error| { + ocr_error( + "use.ocr.install_failed", + format!( + "Failed to open OCR install lock '{}': {error}", + path.display() + ), + ) + })?; + file.lock_exclusive().map_err(|error| { + ocr_error( + "use.ocr.install_failed", + format!( + "Failed to acquire OCR install lock '{}': {error}", + path.display() + ), + ) + })?; + Ok(InstallLock { _file: file }) + }) + .await + .map_err(|error| { + ocr_error( + "use.ocr.install_failed", + format!("OCR install lock task failed: {error}"), + ) + })? +} + +async fn create_stage(root: &Path) -> UseResult { + static NEXT_STAGE: AtomicU64 = AtomicU64::new(1); + for _ in 0..32 { + let path = root.join(format!( + "{STAGE_PREFIX}{}-{}", + std::process::id(), + NEXT_STAGE.fetch_add(1, Ordering::Relaxed) + )); + match tokio::fs::create_dir(&path).await { + Ok(()) => return Ok(path), + Err(error) if error.kind() == std::io::ErrorKind::AlreadyExists => {} + Err(error) => { + return Err(ocr_error( + "use.ocr.install_failed", + format!( + "Failed to create OCR staging directory '{}': {error}", + path.display() + ), + )) + } + } + } + Err(ocr_error( + "use.ocr.install_failed", + "Failed to allocate a unique OCR staging directory.", + )) +} + +async fn cleanup_stale(root: &Path) -> UseResult<()> { + let mut entries = tokio::fs::read_dir(root).await.map_err(|error| { + ocr_error( + "use.ocr.install_failed", + format!( + "Failed to inspect OCR data root '{}': {error}", + root.display() + ), + ) + })?; + while let Some(entry) = entries.next_entry().await.map_err(|error| { + ocr_error( + "use.ocr.install_failed", + format!( + "Failed to inspect OCR data root '{}': {error}", + root.display() + ), + ) + })? { + let name = entry.file_name(); + let name = name.to_string_lossy(); + let is_owned_backup = name.starts_with(BACKUP_PREFIX) && owned_install(&entry.path()); + if name.starts_with(STAGE_PREFIX) || is_owned_backup { + tokio::fs::remove_dir_all(entry.path()) + .await + .map_err(|error| { + ocr_error( + "use.ocr.install_failed", + format!("Failed to remove stale OCR staging directory: {error}"), + ) + })?; + } + } + Ok(()) +} + +async fn activate(stage: &Path, target: &Path) -> UseResult<()> { + static NEXT_BACKUP: AtomicU64 = AtomicU64::new(1); + let parent = target.parent().ok_or_else(|| { + ocr_error( + "use.ocr.install_failed", + "OCR install target has no parent directory.", + ) + })?; + let backup = parent.join(format!( + "{BACKUP_PREFIX}{}-{}", + std::process::id(), + NEXT_BACKUP.fetch_add(1, Ordering::Relaxed) + )); + let had_target = tokio::fs::try_exists(target).await.map_err(|error| { + ocr_error( + "use.ocr.install_failed", + format!( + "Failed to inspect OCR install target '{}': {error}", + target.display() + ), + ) + })?; + if had_target { + if !owned_install(target) { + return Err(ocr_error( + "use.ocr.install_target_unowned", + format!( + "Refusing to replace unowned OCR model directory '{}'.", + target.display() + ), + )); + } + tokio::fs::rename(target, &backup).await.map_err(|error| { + ocr_error( + "use.ocr.install_failed", + format!( + "Failed to stage existing OCR install '{}': {error}", + target.display() + ), + ) + })?; + } + if let Err(error) = tokio::fs::rename(stage, target).await { + if had_target { + let _ = tokio::fs::rename(&backup, target).await; + } + return Err(ocr_error( + "use.ocr.install_failed", + format!( + "Failed to activate OCR install '{}': {error}", + target.display() + ), + )); + } + if had_target { + tokio::fs::remove_dir_all(&backup).await.map_err(|error| { + ocr_error( + "use.ocr.install_failed", + format!( + "Activated OCR but failed to remove backup '{}': {error}", + backup.display() + ), + ) + })?; + } + Ok(()) +} + +fn owned_install(path: &Path) -> bool { + let Ok(bytes) = std::fs::read(path.join(RECEIPT_FILE)) else { + return false; + }; + serde_json::from_slice::(&bytes).is_ok_and(|receipt| { + receipt.schema_version == 1 + && receipt.provider == "pp-ocr-v6" + && receipt.model == MODEL_FAMILY + }) +} + +fn archive_error(error: impl std::fmt::Display) -> UseError { + ocr_error( + "use.ocr.archive_invalid", + format!("Failed to extract PP-OCRv6 archive: {error}"), + ) +} + +fn ocr_error(code: &str, message: impl Into) -> UseError { + UseError::new(code, message) +} diff --git a/crates/ocr/src/lib.rs b/crates/ocr/src/lib.rs index a1bd1740..57486889 100644 --- a/crates/ocr/src/lib.rs +++ b/crates/ocr/src/lib.rs @@ -1,18 +1,27 @@ //! Typed optical character recognition for A3S Use. //! //! OCR is a first-party built-in Use domain and remains process-isolated from -//! A3S Code through its standard MCP server. The crate supports a local -//! Tesseract executable and an explicitly configured OpenAI-compatible vision -//! endpoint without silently installing either provider. +//! A3S Code through its standard MCP server. Detection and recognition run +//! locally with the pinned PP-OCRv6_small ONNX models. There is no alternate +//! OCR provider or off-device fallback. +mod assets; pub mod cli; mod client; +mod config; +mod engine; +mod install; pub mod mcp; mod models; -mod provider; +mod postprocess; +mod preprocess; +pub use assets::{ocr_status, OcrInstallSource, OcrRuntimeStatus}; pub use client::OcrClient; +pub use install::{install_ppocr_v6, repair_ppocr_v6, uninstall_managed_ppocr_v6}; pub use mcp::OcrMcpServer; -pub use models::{OcrBlock, OcrBoundingBox, OcrDiagnostic, OcrProviderKind, OcrRequest, OcrResult}; +pub use models::{ + OcrBlock, OcrBoundingBox, OcrDiagnostic, OcrPoint, OcrProviderKind, OcrRequest, OcrResult, +}; pub use a3s_use_core::{Artifact, Readiness, UseError, UseResult}; diff --git a/crates/ocr/src/mcp.rs b/crates/ocr/src/mcp.rs index f2de4ada..1a3db29f 100644 --- a/crates/ocr/src/mcp.rs +++ b/crates/ocr/src/mcp.rs @@ -43,7 +43,7 @@ impl OcrMcpServer { impl OcrMcpServer { #[tool( name = "ocr_doctor", - description = "Inspect OCR provider readiness without reading an image or making a network request", + description = "Inspect local PP-OCRv6 model readiness without reading an image or making a network request", output_schema = rmcp::handler::server::tool::cached_schema_for_type::(), annotations( read_only_hint = true, @@ -58,13 +58,13 @@ impl OcrMcpServer { #[tool( name = "ocr_extract", - description = "Extract text and available layout evidence from one bounded local image; the configured vision provider may send source bytes to its disclosed endpoint", + description = "Extract text, polygons, bounding boxes, and confidence from one bounded local image with PP-OCRv6; source bytes remain on this device", output_schema = rmcp::handler::server::tool::cached_schema_for_type::(), annotations( read_only_hint = true, destructive_hint = false, idempotent_hint = true, - open_world_hint = true + open_world_hint = false ) )] async fn ocr_extract( @@ -88,7 +88,7 @@ impl ServerHandler for OcrMcpServer { website_url: Some("https://github.com/A3S-Lab/Use".to_string()), }, instructions: Some( - "Call ocr_doctor first. Use ocr_extract only for a local image path supplied by the task. A vision provider may send the complete image to its configured endpoint; do not use it without the user's authority. Preserve the source SHA-256 and distinguish OCR text from verified source text." + "Call ocr_doctor first. Use ocr_extract only for a local image path supplied by the task. PP-OCRv6 detection and recognition run locally through ONNX Runtime and never send source bytes off device. Preserve the source SHA-256 and distinguish OCR text from verified source text." .to_string(), ), ..Default::default() @@ -164,7 +164,7 @@ mod tests { .annotations .as_ref() .and_then(|annotations| annotations.open_world_hint), - Some(true) + Some(false) ); } } diff --git a/crates/ocr/src/models.rs b/crates/ocr/src/models.rs index 6cbf6bae..29b41198 100644 --- a/crates/ocr/src/models.rs +++ b/crates/ocr/src/models.rs @@ -6,9 +6,7 @@ use serde::{Deserialize, Serialize}; #[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize, schemars::JsonSchema)] #[serde(rename_all = "kebab-case")] pub enum OcrProviderKind { - Auto, - Tesseract, - Vision, + PpOcrV6, } #[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize, schemars::JsonSchema)] @@ -16,23 +14,16 @@ pub enum OcrProviderKind { pub struct OcrRequest { #[schemars(description = "Local PNG, JPEG, WebP, GIF, BMP, or TIFF image path")] pub path: PathBuf, - #[serde(default)] - #[schemars( - description = "OCR language identifiers; Tesseract values are joined with '+', for example ['eng', 'chi_sim']" - )] - pub languages: Vec, - #[serde(default)] - #[schemars(description = "Optional Tesseract page segmentation mode from 0 through 13")] - pub page_segmentation_mode: Option, - #[serde(default)] - #[schemars(description = "Override the configured provider for this call")] - pub provider: Option, - #[serde(default)] - #[schemars(description = "Optional extraction instruction used only by the vision provider")] - pub prompt: Option, } -#[derive(Debug, Clone, PartialEq, Serialize, Deserialize, schemars::JsonSchema)] +#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize, schemars::JsonSchema)] +#[serde(rename_all = "camelCase")] +pub struct OcrPoint { + pub x: u32, + pub y: u32, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize, schemars::JsonSchema)] #[serde(rename_all = "camelCase")] pub struct OcrBoundingBox { pub x: u32, @@ -46,19 +37,23 @@ pub struct OcrBoundingBox { pub struct OcrBlock { pub page: u32, pub text: String, - #[serde(skip_serializing_if = "Option::is_none")] - pub confidence: Option, - #[serde(skip_serializing_if = "Option::is_none")] - pub bounding_box: Option, + #[schemars(description = "PP-OCRv6 text recognition confidence from 0 through 1")] + pub confidence: f32, + #[schemars(description = "PP-OCRv6 DB text detection confidence from 0 through 1")] + pub detection_confidence: f32, + #[schemars(description = "Four PP-OCRv6 polygon vertices in source-image coordinates")] + pub polygon: [OcrPoint; 4], + pub bounding_box: OcrBoundingBox, } #[derive(Debug, Clone, PartialEq, Serialize, Deserialize, schemars::JsonSchema)] #[serde(rename_all = "camelCase")] pub struct OcrResult { pub provider: OcrProviderKind, + pub engine: String, + pub model: String, #[schemars(with = "OcrArtifactSchema")] pub source: Artifact, - pub languages: Vec, pub text: String, #[serde(default, skip_serializing_if = "Vec::is_empty")] pub blocks: Vec, @@ -74,11 +69,11 @@ pub struct OcrDiagnostic { #[serde(skip_serializing_if = "Option::is_none")] pub provider: Option, #[serde(skip_serializing_if = "Option::is_none")] - pub executable: Option, - #[serde(skip_serializing_if = "Option::is_none")] - pub endpoint: Option, + pub engine: Option, #[serde(skip_serializing_if = "Option::is_none")] pub model: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub model_dir: Option, pub sends_source_off_device: bool, pub message: String, #[serde(default, skip_serializing_if = "Vec::is_empty")] diff --git a/crates/ocr/src/postprocess.rs b/crates/ocr/src/postprocess.rs new file mode 100644 index 00000000..6104ae94 --- /dev/null +++ b/crates/ocr/src/postprocess.rs @@ -0,0 +1,326 @@ +use a3s_use_core::{UseError, UseResult}; +use clipper2::{Centi, EndType, JoinType}; +use image::{GrayImage, Luma}; +use imageproc::contours::find_contours; +use imageproc::geometry::{contour_area, min_area_rect}; +use imageproc::point::Point; + +use crate::config::{DetectionConfig, RecognitionConfig}; + +#[derive(Debug, Clone)] +pub(crate) struct Detection { + pub(crate) polygon: [Point; 4], + pub(crate) confidence: f32, +} + +#[derive(Debug, Clone, PartialEq)] +pub(crate) struct Recognition { + pub(crate) text: String, + pub(crate) confidence: f32, +} + +pub(crate) fn detection_boxes( + output: &[f32], + shape: &[usize], + original_width: u32, + original_height: u32, + config: &DetectionConfig, +) -> UseResult> { + if shape.len() != 4 || shape[0] != 1 || shape[1] != 1 { + return Err(output_error(format!( + "PP-OCRv6 detection output shape must be [1, 1, H, W], found {shape:?}." + ))); + } + let height = shape[2]; + let width = shape[3]; + let map_len = height + .checked_mul(width) + .ok_or_else(|| output_error("PP-OCRv6 detection output dimensions overflowed."))?; + if width == 0 || height == 0 || output.len() != map_len { + return Err(output_error( + "PP-OCRv6 detection output length does not match its shape.", + )); + } + let width_u32 = u32::try_from(width) + .map_err(|_| output_error("PP-OCRv6 detection output width is too large."))?; + let height_u32 = u32::try_from(height) + .map_err(|_| output_error("PP-OCRv6 detection output height is too large."))?; + let mask = GrayImage::from_fn(width_u32, height_u32, |x, y| { + let index = y as usize * width + x as usize; + Luma([if output[index] > config.threshold { + 255 + } else { + 0 + }]) + }); + + let mut detections = Vec::new(); + for contour in find_contours::(&mask) + .into_iter() + .take(config.max_candidates) + { + if contour.points.len() < 3 { + continue; + } + let mini = order_points(min_area_rect(&contour.points)); + if minimum_side(&mini) < 3.0 { + continue; + } + let score = box_score(output, width, height, &mini); + if score < config.box_threshold { + continue; + } + let area = contour_area(&mini); + let perimeter = polygon_perimeter(&mini); + if !area.is_finite() || !perimeter.is_finite() || perimeter <= f64::EPSILON { + continue; + } + let distance = area * f64::from(config.unclip_ratio) / perimeter; + let path = mini + .iter() + .map(|point| (f64::from(point.x), f64::from(point.y))) + .collect::>(); + let inflated: Vec> = + clipper2::inflate::(path, distance, JoinType::Round, EndType::Polygon, 2.0) + .into(); + if inflated.len() != 1 || inflated[0].len() < 3 { + continue; + } + let inflated = inflated[0] + .iter() + .filter(|(x, y)| x.is_finite() && y.is_finite()) + .map(|(x, y)| { + Point::new( + x.round().clamp(f64::from(i32::MIN), f64::from(i32::MAX)) as i32, + y.round().clamp(f64::from(i32::MIN), f64::from(i32::MAX)) as i32, + ) + }) + .collect::>(); + if inflated.len() < 3 { + continue; + } + let expanded = order_points(min_area_rect(&inflated)); + if minimum_side(&expanded) < 5.0 { + continue; + } + let polygon = expanded.map(|point| { + Point::new( + (point.x as f32 / width as f32 * original_width as f32) + .round() + .clamp(0.0, original_width.saturating_sub(1) as f32), + (point.y as f32 / height as f32 * original_height as f32) + .round() + .clamp(0.0, original_height.saturating_sub(1) as f32), + ) + }); + detections.push(Detection { + polygon, + confidence: score.clamp(0.0, 1.0), + }); + } + sort_reading_order(&mut detections); + Ok(detections) +} + +pub(crate) fn decode_ctc( + output: &[f32], + shape: &[usize], + config: &RecognitionConfig, +) -> UseResult { + if shape.len() != 3 || shape[0] != 1 || shape[1] == 0 || shape[2] == 0 { + return Err(output_error(format!( + "PP-OCRv6 recognition output shape must be [1, T, C], found {shape:?}." + ))); + } + let timesteps = shape[1]; + let classes = shape[2]; + let expected_classes = config.characters.len() + 2; + if classes != expected_classes { + return Err(output_error(format!( + "PP-OCRv6 recognition class count is {classes}, but the model dictionary requires {expected_classes}." + ))); + } + if output.len() != timesteps.saturating_mul(classes) { + return Err(output_error( + "PP-OCRv6 recognition output length does not match its shape.", + )); + } + + let mut text = String::new(); + let mut confidence = 0.0_f32; + let mut selected = 0_usize; + let mut previous = usize::MAX; + for timestep in 0..timesteps { + let row = &output[timestep * classes..(timestep + 1) * classes]; + let (index, score) = row + .iter() + .copied() + .enumerate() + .max_by(|left, right| left.1.total_cmp(&right.1)) + .ok_or_else(|| output_error("PP-OCRv6 recognition output row is empty."))?; + if index != 0 && index != previous { + if index == config.characters.len() + 1 { + text.push(' '); + } else if let Some(character) = config.characters.get(index - 1) { + text.push_str(character); + } + confidence += score; + selected += 1; + } + previous = index; + } + Ok(Recognition { + text, + confidence: if selected == 0 { + 0.0 + } else { + (confidence / selected as f32).clamp(0.0, 1.0) + }, + }) +} + +fn box_score(output: &[f32], width: usize, height: usize, polygon: &[Point; 4]) -> f32 { + let min_x = polygon + .iter() + .map(|point| point.x) + .min() + .unwrap_or(0) + .clamp(0, width.saturating_sub(1) as i32) as usize; + let max_x = polygon + .iter() + .map(|point| point.x) + .max() + .unwrap_or(0) + .clamp(0, width.saturating_sub(1) as i32) as usize; + let min_y = polygon + .iter() + .map(|point| point.y) + .min() + .unwrap_or(0) + .clamp(0, height.saturating_sub(1) as i32) as usize; + let max_y = polygon + .iter() + .map(|point| point.y) + .max() + .unwrap_or(0) + .clamp(0, height.saturating_sub(1) as i32) as usize; + let polygon = polygon.map(|point| Point::new(point.x as f32, point.y as f32)); + let mut sum = 0.0_f32; + let mut count = 0_usize; + for y in min_y..=max_y { + for x in min_x..=max_x { + if point_in_convex_polygon(Point::new(x as f32 + 0.5, y as f32 + 0.5), &polygon) { + sum += output[y * width + x]; + count += 1; + } + } + } + if count == 0 { + 0.0 + } else { + sum / count as f32 + } +} + +fn point_in_convex_polygon(point: Point, polygon: &[Point; 4]) -> bool { + let mut sign = 0_i8; + for index in 0..4 { + let start = polygon[index]; + let end = polygon[(index + 1) % 4]; + let cross = + (end.x - start.x) * (point.y - start.y) - (end.y - start.y) * (point.x - start.x); + if cross.abs() <= f32::EPSILON { + continue; + } + let current = if cross > 0.0 { 1 } else { -1 }; + if sign != 0 && sign != current { + return false; + } + sign = current; + } + true +} + +fn minimum_side(points: &[Point; 4]) -> f64 { + (0..4) + .map(|index| distance(points[index], points[(index + 1) % 4])) + .fold(f64::INFINITY, f64::min) +} + +fn polygon_perimeter(points: &[Point; 4]) -> f64 { + (0..4) + .map(|index| distance(points[index], points[(index + 1) % 4])) + .sum() +} + +fn distance(left: Point, right: Point) -> f64 { + let x = f64::from(left.x - right.x); + let y = f64::from(left.y - right.y); + x.hypot(y) +} + +fn order_points(mut points: [Point; 4]) -> [Point; 4] { + points.sort_by(|left, right| left.x.cmp(&right.x).then(left.y.cmp(&right.y))); + let (top_left, bottom_left) = if points[0].y <= points[1].y { + (points[0], points[1]) + } else { + (points[1], points[0]) + }; + let (top_right, bottom_right) = if points[2].y <= points[3].y { + (points[2], points[3]) + } else { + (points[3], points[2]) + }; + [top_left, top_right, bottom_right, bottom_left] +} + +fn sort_reading_order(detections: &mut [Detection]) { + detections.sort_by(|left, right| { + left.polygon[0] + .y + .total_cmp(&right.polygon[0].y) + .then_with(|| left.polygon[0].x.total_cmp(&right.polygon[0].x)) + }); + for index in 1..detections.len() { + let mut cursor = index; + while cursor > 0 { + let current = detections[cursor].polygon[0]; + let previous = detections[cursor - 1].polygon[0]; + if (current.y - previous.y).abs() < 10.0 && current.x < previous.x { + detections.swap(cursor, cursor - 1); + cursor -= 1; + } else { + break; + } + } + } +} + +fn output_error(message: impl Into) -> UseError { + UseError::new("use.ocr.provider_output_invalid", message) +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn ctc_decoder_removes_blanks_and_repeated_classes() { + let config = RecognitionConfig { + channels: 3, + height: 48, + default_width: 320, + characters: vec!["A".to_string(), "B".to_string()], + }; + let output = [ + 0.9, 0.1, 0.0, 0.0, // blank + 0.1, 0.8, 0.1, 0.0, // A + 0.9, 0.1, 0.0, 0.0, // blank + 0.1, 0.8, 0.1, 0.0, // A + 0.1, 0.1, 0.8, 0.0, // B + ]; + let result = decode_ctc(&output, &[1, 5, 4], &config).unwrap(); + assert_eq!(result.text, "AAB"); + assert!((result.confidence - 0.8).abs() < f32::EPSILON); + } +} diff --git a/crates/ocr/src/preprocess.rs b/crates/ocr/src/preprocess.rs new file mode 100644 index 00000000..4cdcd9f7 --- /dev/null +++ b/crates/ocr/src/preprocess.rs @@ -0,0 +1,192 @@ +use std::io::Cursor; + +use a3s_use_core::{UseError, UseResult}; +use image::imageops::FilterType; +use image::{DynamicImage, ImageReader, Limits, RgbImage}; + +use crate::config::{DetectionConfig, RecognitionConfig}; + +const MAX_IMAGE_SIDE: u32 = 16_384; +const MAX_DECODED_BYTES: u64 = 256 * 1024 * 1024; +const DETECTION_MIN_SIDE: u32 = 736; +const DETECTION_MAX_SIDE: u32 = 4_000; +const RECOGNITION_MAX_WIDTH: u32 = 3_200; + +pub(crate) struct DetectionInput { + pub(crate) data: Vec, + pub(crate) shape: [usize; 4], + pub(crate) original_width: u32, + pub(crate) original_height: u32, +} + +pub(crate) struct RecognitionInput { + pub(crate) data: Vec, + pub(crate) shape: [usize; 4], +} + +pub(crate) fn decode_image(bytes: &[u8]) -> UseResult { + let cursor = Cursor::new(bytes); + let mut reader = ImageReader::new(cursor) + .with_guessed_format() + .map_err(|error| image_error(format!("Failed to detect OCR image format: {error}")))?; + let mut limits = Limits::default(); + limits.max_image_width = Some(MAX_IMAGE_SIDE); + limits.max_image_height = Some(MAX_IMAGE_SIDE); + limits.max_alloc = Some(MAX_DECODED_BYTES); + reader.limits(limits); + let image = reader + .decode() + .map_err(|error| image_error(format!("Failed to decode OCR image: {error}")))?; + let width = image.width(); + let height = image.height(); + if width == 0 + || height == 0 + || u64::from(width) + .checked_mul(u64::from(height)) + .and_then(|pixels| pixels.checked_mul(4)) + .is_none_or(|bytes| bytes > MAX_DECODED_BYTES) + { + return Err(image_error( + "Decoded OCR image dimensions exceed the 256 MiB pixel limit.", + )); + } + Ok(image.to_rgb8()) +} + +pub(crate) fn detection_input( + image: &RgbImage, + config: &DetectionConfig, +) -> UseResult { + let original_width = image.width(); + let original_height = image.height(); + let (width, height) = detection_dimensions(original_width, original_height)?; + let resized = if width == original_width && height == original_height { + image.clone() + } else { + DynamicImage::ImageRgb8(image.clone()) + .resize_exact(width, height, FilterType::Triangle) + .to_rgb8() + }; + let plane = usize::try_from(u64::from(width) * u64::from(height)) + .map_err(|_| image_error("Detection tensor dimensions overflowed."))?; + let mut data = vec![0.0_f32; plane * 3]; + for (index, pixel) in resized.pixels().enumerate() { + let channels = [pixel[2], pixel[1], pixel[0]]; + for channel in 0..3 { + data[channel * plane + index] = (f32::from(channels[channel]) * config.scale + - config.mean[channel]) + / config.std[channel]; + } + } + Ok(DetectionInput { + data, + shape: [1, 3, height as usize, width as usize], + original_width, + original_height, + }) +} + +pub(crate) fn recognition_input( + images: &[RgbImage], + config: &RecognitionConfig, +) -> UseResult { + if images.is_empty() || images.len() > 8 { + return Err(image_error( + "PP-OCRv6 recognition batches must contain from 1 through 8 text crops.", + )); + } + if images + .iter() + .any(|image| image.width() == 0 || image.height() == 0) + { + return Err(image_error("PP-OCRv6 text crop has zero width or height.")); + } + let model_height = u32::try_from(config.height) + .map_err(|_| image_error("Recognition model height is invalid."))?; + let default_width = u32::try_from(config.default_width) + .map_err(|_| image_error("Recognition model width is invalid."))?; + let resized_widths = images + .iter() + .map(|image| { + ((f64::from(model_height) * f64::from(image.width()) / f64::from(image.height())).ceil() + as u32) + .clamp(1, RECOGNITION_MAX_WIDTH) + }) + .collect::>(); + let widest = resized_widths.iter().copied().max().unwrap_or(1); + let canvas_width = default_width.max(widest).min(RECOGNITION_MAX_WIDTH); + let target_plane = usize::try_from(u64::from(canvas_width) * u64::from(model_height)) + .map_err(|_| image_error("Recognition tensor dimensions overflowed."))?; + let batch_stride = config + .channels + .checked_mul(target_plane) + .ok_or_else(|| image_error("Recognition tensor dimensions overflowed."))?; + let mut data = vec![0.0_f32; images.len() * batch_stride]; + for (batch, (image, resized_width)) in images.iter().zip(resized_widths).enumerate() { + let resized = DynamicImage::ImageRgb8(image.clone()) + .resize_exact(resized_width, model_height, FilterType::Triangle) + .to_rgb8(); + for y in 0..model_height { + for x in 0..resized_width { + let pixel = resized.get_pixel(x, y); + let target = y as usize * canvas_width as usize + x as usize; + let channels = [pixel[2], pixel[1], pixel[0]]; + for channel in 0..config.channels { + data[batch * batch_stride + channel * target_plane + target] = + f32::from(channels[channel]) / 127.5 - 1.0; + } + } + } + } + Ok(RecognitionInput { + data, + shape: [ + images.len(), + config.channels, + config.height, + canvas_width as usize, + ], + }) +} + +fn detection_dimensions(width: u32, height: u32) -> UseResult<(u32, u32)> { + if width == 0 || height == 0 { + return Err(image_error("OCR image has zero width or height.")); + } + let mut ratio = 1.0_f64; + let min_side = width.min(height); + let max_side = width.max(height); + if min_side < DETECTION_MIN_SIDE { + ratio = f64::from(DETECTION_MIN_SIDE) / f64::from(min_side); + } + if f64::from(max_side) * ratio > f64::from(DETECTION_MAX_SIDE) { + ratio = f64::from(DETECTION_MAX_SIDE) / f64::from(max_side); + } + let resized_width = round_stride(f64::from(width) * ratio, 32); + let resized_height = round_stride(f64::from(height) * ratio, 32); + Ok((resized_width, resized_height)) +} + +fn round_stride(value: f64, stride: u32) -> u32 { + let rounded = (value / f64::from(stride)).round_ties_even() as u32 * stride; + rounded.max(stride) +} + +fn image_error(message: impl Into) -> UseError { + UseError::new("use.ocr.image_invalid", message) +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn detection_dimensions_are_bounded_stride_multiples() { + assert_eq!(detection_dimensions(10, 20).unwrap(), (736, 1_472)); + assert_eq!(detection_dimensions(4_000, 1_000).unwrap(), (4_000, 992)); + assert_eq!( + detection_dimensions(20_000, 10_000).unwrap(), + (4_000, 1_984) + ); + } +} diff --git a/crates/ocr/tests/ppocr_v6_contract.rs b/crates/ocr/tests/ppocr_v6_contract.rs new file mode 100644 index 00000000..8f27ccd2 --- /dev/null +++ b/crates/ocr/tests/ppocr_v6_contract.rs @@ -0,0 +1,19 @@ +use std::path::PathBuf; + +use a3s_use_ocr::{OcrProviderKind, OcrRequest}; + +#[test] +fn public_contract_names_only_pp_ocr_v6() { + assert_eq!( + serde_json::to_value(OcrProviderKind::PpOcrV6).unwrap(), + serde_json::json!("pp-ocr-v6") + ); + + let request = OcrRequest { + path: PathBuf::from("scan.png"), + }; + assert_eq!( + serde_json::to_value(request).unwrap(), + serde_json::json!({ "path": "scan.png" }) + ); +} diff --git a/docs/architecture.md b/docs/architecture.md index c165d7b5..92d7603d 100644 --- a/docs/architecture.md +++ b/docs/architecture.md @@ -41,10 +41,10 @@ does not define JSON-RPC methods or convert surfaces implicitly. Use build. The release packages its content-bound Skill and exposes the native CLI plus standard stdio MCP without a separate extension install. The process accepts bounded local image files and binds every result to the canonical -source digest. It uses only an explicitly present Tesseract executable or -configured vision endpoint; it never installs a provider silently. The vision -diagnostic discloses off-device image transfer, and its MCP extraction tool -carries conservative open-world annotations. +source digest. It runs the pinned `PP-OCRv6_small` detection and recognition +models locally through ONNX Runtime, without Python, PaddlePaddle, a remote OCR +endpoint, or an alternate backend. Model installation and repair are explicit +`use/ocr` component operations. Both MCP tools are closed-world and read-only. ## Hot-plug registry @@ -89,8 +89,8 @@ compatibility provider. For resident hosts, `use/office` targets the built-in `office-native` MCP server and is ready independently of OfficeCLI. A discovered OfficeCLI provider is projected separately as `use/office-compat`, targeting the standard compatibility server without carrying the native Skill. The -`use/ocr` route targets `ocr-native`; provider readiness remains explicit and -never triggers a silent Tesseract or vision-provider install. +`use/ocr` route targets `ocr-native`; model readiness remains explicit and +never triggers a silent install. The projection contains content-bound Skill references and an MCP launch target, never executable extension code or a generic action payload. Consumers still @@ -495,11 +495,11 @@ Implemented: 15. A packaged first-party `a3s-use-office` Skill with progressive Word/Spreadsheet/Presentation/MCP references, bounded local discovery, release-archive smoke checks, and content-bound capability projection. -16. A first-party built-in OCR route with typed provider diagnostics, bounded - image admission, source SHA-256 evidence, local Tesseract and explicit - vision adapters, standard MCP annotations/output schemas, and a - release-packaged content-bound Skill that projects to `mcp__use_ocr__*` in - A3S Code. +16. A first-party built-in OCR route with typed PP-OCRv6 diagnostics, bounded + image admission, source SHA-256 evidence, native ONNX Runtime detection and + recognition, standard closed-world MCP annotations/output schemas, pinned + release models, and a content-bound Skill that projects to + `mcp__use_ocr__*` in A3S Code. Next: diff --git a/src/cli.rs b/src/cli.rs index 23bad00f..e0875d02 100644 --- a/src/cli.rs +++ b/src/cli.rs @@ -119,7 +119,7 @@ fn help() -> CommandOutput { " a3s-use office native get|query|view|watch|raw|raw-set|dump|merge|validate|create|add|add-part|set|sort|remove|move|copy|swap|insert-rows|delete-rows|insert-columns|delete-columns|rename-sheet|move-sheet|copy-sheet|batch [args] [--json]\n", " a3s-use office \n", " a3s-use ocr doctor [--json]\n", - " a3s-use ocr extract [--language ] [--provider ] [--json]\n", + " a3s-use ocr extract [--json]\n", " a3s-use extension list|inspect|doctor [args] [--json]\n", " a3s-use extension enable [--json]\n", " a3s-use extension disable [--timeout-ms ] [--json]\n", @@ -416,6 +416,36 @@ async fn component_install(args: &[String]) -> UseResult { )); } } + if matches!(id, "ocr" | "use/ocr") { + #[cfg(feature = "ocr")] + { + if option_argument(args, "--from")?.is_some() { + return Err(usage_error("--from is valid only for external extensions")); + } + let force = args.iter().any(|argument| argument == "--force"); + let previous = a3s_use_ocr::ocr_status(); + let status = a3s_use_ocr::install_ppocr_v6(force).await?; + let changed = force + || !previous.available + || previous.model_dir != status.model_dir + || previous.source != status.source; + let diagnostic = ocr_diagnostic(); + return Ok(CommandOutput::success( + format!( + "Local PP-OCRv6 model bundle is ready at {}.", + status.model_dir.as_ref().map_or_else( + || "an unknown path".to_string(), + |path| path.display().to_string() + ) + ), + serde_json::json!({ + "component": component_value(id, &diagnostic), + "changed": changed, + "runtime": status + }), + )); + } + } if let Some(diagnostic) = builtin_diagnostic(id) { if option_argument(args, "--from")?.is_some() { return Err(usage_error("--from is valid only for external extensions")); @@ -513,6 +543,24 @@ async fn component_uninstall(id: &str) -> UseResult { )); } } + if matches!(id, "ocr" | "use/ocr") { + #[cfg(feature = "ocr")] + { + let changed = a3s_use_ocr::uninstall_managed_ppocr_v6().await?; + return Ok(CommandOutput::success( + if changed { + "Removed A3S-managed PP-OCRv6 model files." + } else { + "No A3S-managed PP-OCRv6 model files are installed." + }, + serde_json::json!({ + "component": id, + "changed": changed, + "builtInCommandPreserved": true + }), + )); + } + } if matches!( id, "browser" | "use/browser" | "office" | "use/office" | "ocr" | "use/ocr" @@ -896,6 +944,8 @@ fn builtin_presence(id: &str) -> &'static str { ), #[cfg(feature = "office")] "office" | "use/office" => office_presence(a3s_use_office::office_status().source), + #[cfg(feature = "ocr")] + "ocr" | "use/ocr" => ocr_presence(a3s_use_ocr::ocr_status().source), _ => "external", } } @@ -921,6 +971,16 @@ fn office_presence(source: a3s_use_office::OfficeInstallSource) -> &'static str } } +#[cfg(feature = "ocr")] +fn ocr_presence(source: a3s_use_ocr::OcrInstallSource) -> &'static str { + match source { + a3s_use_ocr::OcrInstallSource::Environment => "external", + a3s_use_ocr::OcrInstallSource::Packaged => "packaged", + a3s_use_ocr::OcrInstallSource::Managed => "managed", + a3s_use_ocr::OcrInstallSource::Missing => "missing", + } +} + fn builtin_diagnostic(id: &str) -> Option { match id { "browser" | "use/browser" => Some(browser_diagnostic()), diff --git a/src/cli_tests.rs b/src/cli_tests.rs index e60610cd..23ff377b 100644 --- a/src/cli_tests.rs +++ b/src/cli_tests.rs @@ -157,3 +157,14 @@ fn office_component_presence_preserves_runtime_ownership() { assert_eq!(office_presence(OfficeInstallSource::Managed), "managed"); assert_eq!(office_presence(OfficeInstallSource::Missing), "missing"); } + +#[cfg(feature = "ocr")] +#[test] +fn ocr_component_presence_preserves_model_ownership() { + use a3s_use_ocr::OcrInstallSource; + + assert_eq!(ocr_presence(OcrInstallSource::Environment), "external"); + assert_eq!(ocr_presence(OcrInstallSource::Packaged), "packaged"); + assert_eq!(ocr_presence(OcrInstallSource::Managed), "managed"); + assert_eq!(ocr_presence(OcrInstallSource::Missing), "missing"); +} diff --git a/src/ocr_builtin.rs b/src/ocr_builtin.rs index 127004c2..0e983c18 100644 --- a/src/ocr_builtin.rs +++ b/src/ocr_builtin.rs @@ -14,7 +14,7 @@ pub(crate) fn diagnostic() -> DomainDiagnostic { readiness: diagnostic.readiness, provider: diagnostic.provider.map(provider_name).map(str::to_string), version: None, - path: diagnostic.executable, + path: diagnostic.model_dir, message: diagnostic.message, suggestions: diagnostic.suggestions, } @@ -65,9 +65,7 @@ pub(crate) async fn primary_skill_surface() -> Option<(PathBuf, PathBuf)> { fn provider_name(provider: OcrProviderKind) -> &'static str { match provider { - OcrProviderKind::Auto => "auto", - OcrProviderKind::Tesseract => "tesseract", - OcrProviderKind::Vision => "vision", + OcrProviderKind::PpOcrV6 => "pp-ocr-v6", } } From 7389266c2b0aedce1b922c99388965611fb0028a Mon Sep 17 00:00:00 2001 From: RoyLin Date: Sun, 19 Jul 2026 10:43:10 +0800 Subject: [PATCH 5/9] feat(office): add native spreadsheet formulas and import --- Cargo.toml | 2 +- README.md | 173 +++- crates/office/skills/a3s-use-office/SKILL.md | 51 +- .../skills/a3s-use-office/references/mcp.md | 121 ++- .../a3s-use-office/references/spreadsheet.md | 165 +++- crates/office/src/editor.rs | 97 +- crates/office/src/editor/spreadsheet.rs | 58 +- .../office/src/editor/spreadsheet/formula.rs | 169 ++++ .../editor/spreadsheet/formula/planning.rs | 337 +++++++ .../src/editor/spreadsheet/formula/write.rs | 354 +++++++ .../office/src/editor/spreadsheet/import.rs | 514 ++++++++++ .../src/editor/spreadsheet/import/parse.rs | 435 +++++++++ crates/office/src/editor/spreadsheet/style.rs | 20 +- crates/office/src/editor/spreadsheet/table.rs | 127 ++- .../src/editor/spreadsheet/table/formula.rs | 299 ++++++ crates/office/src/editor/spreadsheet/view.rs | 300 ++++++ crates/office/src/editor/types.rs | 22 + .../src/editor/types/spreadsheet_import.rs | 109 +++ .../src/editor/types/spreadsheet_view.rs | 130 +++ crates/office/src/issues.rs | 4 +- crates/office/src/lib.rs | 43 +- crates/office/src/replay.rs | 252 ++++- crates/office/src/semantic/mod.rs | 2 + crates/office/src/semantic/selector.rs | 1 + crates/office/src/semantic/spreadsheet.rs | 12 + .../office/src/semantic/spreadsheet/view.rs | 90 ++ crates/office/src/semantic_tests.rs | 29 + crates/office/src/spreadsheet_formula.rs | 199 +++- crates/office/src/spreadsheet_formula/ast.rs | 245 +++++ .../src/spreadsheet_formula/evaluate.rs | 415 ++++++++ .../spreadsheet_formula/evaluate/context.rs | 406 ++++++++ .../spreadsheet_formula/evaluate/function.rs | 177 ++++ .../evaluate/function/aggregate.rs | 392 ++++++++ .../evaluate/function/array.rs | 145 +++ .../evaluate/function/lazy.rs | 170 ++++ .../spreadsheet_formula/evaluate/operators.rs | 243 +++++ .../spreadsheet_formula/evaluate/reference.rs | 615 ++++++++++++ .../src/spreadsheet_formula/evaluate/spill.rs | 191 ++++ .../office/src/spreadsheet_formula/graph.rs | 515 ++++++++++ .../spreadsheet_formula/graph/reference.rs | 582 +++++++++++ .../office/src/spreadsheet_formula/lexer.rs | 823 ++++++++++++++++ .../office/src/spreadsheet_formula/parser.rs | 923 ++++++++++++++++++ .../src/spreadsheet_formula/parser/tests.rs | 106 ++ .../src/spreadsheet_formula/registry.rs | 343 +++++++ .../structured_reference.rs | 516 ++++++++++ .../structured_reference/parser.rs | 270 +++++ .../structured_reference/rewrite.rs | 538 ++++++++++ .../office/src/spreadsheet_formula/value.rs | 58 ++ .../spreadsheet_formula_calculation_tests.rs | 2 + .../calculation.rs | 759 ++++++++++++++ .../recalculation.rs | 547 +++++++++++ .../src/spreadsheet_formula_graph_tests.rs | 226 +++++ crates/office/src/spreadsheet_import_tests.rs | 356 +++++++ docs/native-office.md | 175 +++- src/mcp/office/input.rs | 31 + src/mcp/office/input/spreadsheet_import.rs | 36 + src/mcp/office/input/spreadsheet_view.rs | 19 + src/mcp/office/tests.rs | 92 ++ src/office_native_cli.rs | 8 +- src/office_native_cli/bounded_input.rs | 20 + src/office_native_cli/spreadsheet_formula.rs | 40 + src/office_native_cli/spreadsheet_import.rs | 251 +++++ tests/cli.rs | 25 + tests/office_cell_format_mcp.rs | 36 + tests/office_spreadsheet_formula_cli.rs | 159 +++ tests/office_spreadsheet_formula_mcp.rs | 249 +++++ tests/office_spreadsheet_import_cli.rs | 275 ++++++ tests/office_spreadsheet_import_mcp.rs | 331 +++++++ 68 files changed, 15278 insertions(+), 147 deletions(-) create mode 100644 crates/office/src/editor/spreadsheet/formula.rs create mode 100644 crates/office/src/editor/spreadsheet/formula/planning.rs create mode 100644 crates/office/src/editor/spreadsheet/formula/write.rs create mode 100644 crates/office/src/editor/spreadsheet/import.rs create mode 100644 crates/office/src/editor/spreadsheet/import/parse.rs create mode 100644 crates/office/src/editor/spreadsheet/table/formula.rs create mode 100644 crates/office/src/editor/spreadsheet/view.rs create mode 100644 crates/office/src/editor/types/spreadsheet_import.rs create mode 100644 crates/office/src/editor/types/spreadsheet_view.rs create mode 100644 crates/office/src/semantic/spreadsheet/view.rs create mode 100644 crates/office/src/spreadsheet_formula/ast.rs create mode 100644 crates/office/src/spreadsheet_formula/evaluate.rs create mode 100644 crates/office/src/spreadsheet_formula/evaluate/context.rs create mode 100644 crates/office/src/spreadsheet_formula/evaluate/function.rs create mode 100644 crates/office/src/spreadsheet_formula/evaluate/function/aggregate.rs create mode 100644 crates/office/src/spreadsheet_formula/evaluate/function/array.rs create mode 100644 crates/office/src/spreadsheet_formula/evaluate/function/lazy.rs create mode 100644 crates/office/src/spreadsheet_formula/evaluate/operators.rs create mode 100644 crates/office/src/spreadsheet_formula/evaluate/reference.rs create mode 100644 crates/office/src/spreadsheet_formula/evaluate/spill.rs create mode 100644 crates/office/src/spreadsheet_formula/graph.rs create mode 100644 crates/office/src/spreadsheet_formula/graph/reference.rs create mode 100644 crates/office/src/spreadsheet_formula/lexer.rs create mode 100644 crates/office/src/spreadsheet_formula/parser.rs create mode 100644 crates/office/src/spreadsheet_formula/parser/tests.rs create mode 100644 crates/office/src/spreadsheet_formula/registry.rs create mode 100644 crates/office/src/spreadsheet_formula/structured_reference.rs create mode 100644 crates/office/src/spreadsheet_formula/structured_reference/parser.rs create mode 100644 crates/office/src/spreadsheet_formula/structured_reference/rewrite.rs create mode 100644 crates/office/src/spreadsheet_formula/value.rs create mode 100644 crates/office/src/spreadsheet_formula_calculation_tests.rs create mode 100644 crates/office/src/spreadsheet_formula_calculation_tests/calculation.rs create mode 100644 crates/office/src/spreadsheet_formula_calculation_tests/recalculation.rs create mode 100644 crates/office/src/spreadsheet_formula_graph_tests.rs create mode 100644 crates/office/src/spreadsheet_import_tests.rs create mode 100644 src/mcp/office/input/spreadsheet_import.rs create mode 100644 src/mcp/office/input/spreadsheet_view.rs create mode 100644 src/office_native_cli/spreadsheet_formula.rs create mode 100644 src/office_native_cli/spreadsheet_import.rs create mode 100644 tests/office_spreadsheet_formula_cli.rs create mode 100644 tests/office_spreadsheet_formula_mcp.rs create mode 100644 tests/office_spreadsheet_import_cli.rs create mode 100644 tests/office_spreadsheet_import_mcp.rs diff --git a/Cargo.toml b/Cargo.toml index a136a13b..5ba64787 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -35,7 +35,7 @@ serde_json = "1" sha2 = "0.10" thiserror = "2" tempfile = "3" -tokio = { version = "1", features = ["fs", "io-util", "macros", "net", "rt-multi-thread", "process", "sync", "time"] } +tokio = { version = "1", features = ["fs", "io-std", "io-util", "macros", "net", "rt-multi-thread", "process", "sync", "time"] } tokio-util = "0.7" url = "2" zip = { version = "2", default-features = false, features = ["deflate"] } diff --git a/README.md b/README.md index 9b8376a1..92a82536 100644 --- a/README.md +++ b/README.md @@ -70,6 +70,8 @@ a3s use office native set report.docx '/body/p[1]' --align center --json a3s use office native set workbook.xlsx /Sheet1/A1:C3 --number-format currency --fill FFF2CC --border-all thin --border-color C9B458 --vertical-align center --wrap-text true --json a3s use office native set workbook.xlsx /Sheet1/A1:C1 --text 'Quarter' --bold true --merge-cells true --json a3s use office native sort workbook.xlsx /Sheet1/A1:D100 --key B:desc --key C:asc --header true --case-sensitive false --json +a3s use office native import workbook.xlsx /Sheet1 source.csv --header --start-cell A1 --json +a3s use office native recalculate workbook.xlsx --output calculated.xlsx --json a3s use office native add workbook.xlsx /Sheet1 --type table --name Sales --range F1:H4 --table-column Name --table-column Qty --table-column Price --style medium:4 --json a3s use office native add report.docx '/body/p[1]' --type hyperlink --url https://example.com --display 'Open site' --tooltip 'A3S site' --json a3s use office native add report.docx '/body/p[1]' --type comment --author Alice --initials AL --text 'Please review' --json @@ -115,9 +117,12 @@ Every domain argument accepted by `a3s use ...` can also be passed directly to scoped cross-format literal/regex replacement, typed text formatting, Spreadsheet number/fill/border/alignment formatting, exact merged-cell editing, stable multi-key physical row sorting with persisted sort state, - typed worksheet and table AutoFilters, typed data validation and conditional - formatting, scoped defined names, native Spreadsheet ListObject table - lifecycle, inert hyperlinks, and + bounded CSV/TSV import with typed inference and canonical header + filter/freeze behavior, a bounded formula parser and dependency graph, + explicit native formula recalculation with typed cached values and dynamic + arrays, typed worksheet and table AutoFilters, typed data validation and + conditional formatting, scoped defined names, native Spreadsheet ListObject + table lifecycle, inert hyperlinks, and legacy comments, native PNG/JPEG/GIF embedding, cross-format template merge, deterministic bounded all-format annotated views, all-format HTML/SVG semantic previews, @@ -361,9 +366,10 @@ protects `_xlnm.*` and `Slicer_*` names owned by other Office features. Semantic get/query, ordinary typed remove, batch, exact replay, CLI, Rust, and standard MCP share the same value. Strict/transitional SpreadsheetML and unknown defined-name attributes are retained; unknown collection or child -content fails closed when it cannot be preserved. This is defined-name -lifecycle support, not formula evaluation, external-link authoring, or complete -Spreadsheet parity. +content fails closed when it cannot be preserved. This contract owns +defined-name lifecycle, not external-link authoring or complete Spreadsheet +parity; supported names participate when an explicit native cell-formula +recalculation references them. Native Spreadsheet AutoFilters use a closed typed contract shared by worksheet filters and ListObject tables. One value owns a normalized rectangular A1 @@ -405,6 +411,63 @@ ignored errors, and supported drawing anchors move with their records; chart caches are cleared and the worksheet used dimension is recomputed after the physical change. +Native Spreadsheet delimited import accepts bounded UTF-8 CSV or TSV from a +regular file or stdin and writes it into an existing worksheet from an explicit +A1 start cell. The parser supports a leading BOM, CRLF, quoted delimiters, +embedded newlines, and doubled quotes; malformed quote state fails atomically. +One request is limited to 8 MiB and a 100,000-cell rectangular extent. Explicit +empty fields clear existing target cells, while missing trailing fields in a +ragged row leave those cells unchanged. + +Typed inference stores formulas, finite numbers, booleans, ISO dates/times, and +text without a Python, Node.js, OfficeCLI, or spreadsheet-application runtime. +Dates honor the workbook's 1900/1904 system and receive a native date number +format. Inferred formulas pass the same bounded native syntax parser as direct +cell writes; malformed expressions fail the complete import. Import does not +implicitly calculate formulas; run the explicit native recalculation command +or mutation when fresh cached results are required. Header mode atomically +installs the worksheet AutoFilter and a canonical frozen pane below the header. +Frozen pane state is readable at `/Sheet/freeze`, is set through the typed +batch/Rust/MCP contract, and is removed through the ordinary typed `remove` +mutation. Strict/transitional SpreadsheetML and unknown view content are +preserved; unsupported pane content reports `nativeMutable=false` and fails +closed on mutation. + +Native Spreadsheet formula calculation builds a deterministic bounded +dependency graph across worksheets, ranges, spills, and workbook- or +worksheet-scoped names. The closed built-in registry implements +`SUM`, `AVERAGE`, `MIN`, `MAX`, `COUNT`, `COUNTA`, `ABS`, `SQRT`, `POWER`, +`MOD`, `ROUND`, `IF`, `IFERROR`, `AND`, `OR`, `NOT`, `CONCAT`, +`CONCATENATE`, `ROW`, `COLUMN`, `SEQUENCE`, `TRANSPOSE`, `PI`, and `NA`. +Operators and typed blank, number, text, boolean, and Spreadsheet error values +participate in scalar or rectangular array calculation. + +The read-only Rust calculation API leaves package bytes unchanged. The editor, +versioned batch, replay, CLI `office native recalculate`, and standard MCP +`recalculate-spreadsheet-formulas` mutation atomically write typed OOXML +caches, canonical array anchors, spill children, and calculated-workbook +metadata. Spill children are read-only; edit or remove their formula anchor. +Exact replay accepts canonical formula storage and natively cached array +anchors. It fails closed for physical distinctions typed mutations cannot +reproduce, including explicit `t="normal"` storage and uncached or malformed +array anchors. +ListObject structured references resolve table names or display names. +`Sales[Qty]` and `Sales[[Qty]:[Price]]` select data rows; `#All`, `#Data`, +`#Headers`, and `#Totals` select structural rows; and `Sales[@Qty]`, +`Sales[[#This Row],[Qty]]`, or table-local `[@Qty]` select the current data row. +Table-local forms require the formula cell to be inside the inferred table. +Missing tables, columns, or requested structural rows, disjoint columns, +non-canonical forms, cycles, unsupported or qualified functions, and +external-workbook reads fail with stable typed errors and roll back the whole +batch. The engine never fetches an external workbook or falls back to a shell +or script runtime. One formula is limited to 8,192 characters, depth +128 (including nested named-reference resolution), and 8,192 AST nodes; one +reference value to 100,000 areas; one graph to 100,000 formulas, 1,000,000 +edges, and 1,000,000 formula-cell reference visits; one materialized array or +function call to 100,000 cells; and one text result to 1 MiB. +One calculation pass is also limited to 100,000 cumulative spill children and +200,000 OOXML cell writes, plus 8 MiB of cumulative text-result bytes. + Native Spreadsheet tables use a separate closed ListObject contract. Add and set own the workbook-wide `name`, optional distinct `displayName`, final rectangular A1 range, one exact column identity per range column, header/totals @@ -412,7 +475,17 @@ row state, typed filter criteria, built-in light/medium/dark style identity, and first/last-column plus row/column-stripe flags; ordinary typed `remove` owns deletion. When a header is enabled, its names are stamped into the first row and the table-owned AutoFilter range excludes an enabled totals row. The -editor rejects +`set` lifecycle keeps common structured references consistent: changed table +names/display names and position-mapped column identities are rewritten across +cell formulas, defined names, conditional formats, data validations, charts, +and table formulas without touching string literals or external-workbook +references. Table-local references are rewritten only when their ListObject +context is provable. Unsafe local-reference geometry changes fail with +`use.office.spreadsheet_table_formula_rewrite_unsupported`, and `remove` +fails with `use.office.spreadsheet_table_referenced` while a live structured +reference still targets the table. All checks and rewrites share the table +mutation's atomic rollback boundary. +The editor rejects Excel-identifier and A1/R1C1 name errors, case-insensitive table/defined-name collisions, duplicate columns, missing data rows, table/merge/worksheet-AutoFilter overlap, and unsafe relationship graphs. @@ -495,10 +568,11 @@ Browser contract. The explicit `office native` CLI exposes in-process blank creation, reads, typed add/set/remove/move/copy/swap, scoped literal/regex replacement, rich-text, exact Spreadsheet merged-cell, stable Spreadsheet physical sorting -with persisted sort state, worksheet/table AutoFilter, -data-validation, conditional-format, defined-name, ListObject table, hyperlink, -and legacy-comment -operations, constrained raw XML access, +with persisted sort state, bounded CSV/TSV import with typed inference and +header filter/freeze behavior, explicit Spreadsheet formula recalculation, +worksheet/table AutoFilter, data-validation, conditional-format, defined-name, +ListObject table, hyperlink, and legacy-comment operations, constrained raw +XML access, known typed part carriers, exact replay artifacts for a constrained canonical subset, visible PNG/JPEG/GIF pictures, and atomic mutation batches, plus dependency-free template merge and semantic rendering today. HTML and SVG are @@ -592,6 +666,12 @@ a3s use office native set workbook.xlsx /Sheet1/E1 --border-diagonal slant-dash- a3s use office native set workbook.xlsx /Sheet1/A1:C1 --text 'Quarter' --bold true --merge-cells true --json a3s use office native set workbook.xlsx /Sheet1/A1:C1 --merge-cells false --json +# Import a bounded UTF-8 CSV/TSV source. Header mode also installs the +# worksheet AutoFilter and canonical frozen pane in the same transaction. +a3s use office native import workbook.xlsx /Sheet1 source.csv --header --start-cell A1 --json +a3s use office native import workbook.xlsx /Sheet1 --stdin --format tsv --output imported.xlsx --json +a3s use office native get workbook.xlsx /Sheet1/freeze --json + # Add, inspect, replace, clear, and remove one worksheet AutoFilter. Each # --filter is a strict JSON object with a zero-based column and typed criteria. a3s use office native add workbook.xlsx /Sheet1 --type auto-filter --range A1:C20 --filter '{"column":0,"criteria":{"type":"values","values":["Open","Closed"],"includeBlanks":true}}' --filter '{"column":2,"criteria":{"type":"greater-than","value":"100"}}' --json @@ -669,10 +749,12 @@ a3s use office native add deck.pptx '/slide[1]' --type comment --author Alice -- a3s use office native query deck.pptx comment --json a3s use office native remove workbook.xlsx /Sheet1/B2/comment --json -# Preserve Spreadsheet value types; formula storage requests application recalculation. +# Preserve Spreadsheet value types. Formula writes validate and store the +# expression; explicit recalculation computes and writes cached values. a3s use office native set workbook.xlsx /Sheet1/A1 --number 42.5 --json a3s use office native set workbook.xlsx /Sheet1/B1 --boolean true --json a3s use office native set workbook.xlsx /Sheet1/C1 --formula 'SUM(A1:B1)' --json +a3s use office native recalculate workbook.xlsx --output calculated.xlsx --json # Set or remove a bounded rectangular range atomically. a3s use office native set workbook.xlsx /Sheet1/A2:C4 --number 0 --json @@ -851,8 +933,9 @@ and 10,000 mutations. The version 1 mutation set is `replace-text`, `set-text`, `set-data-validation`, `add-conditional-format`, `set-conditional-format`, `add-named-range`, `set-named-range`, `add-spreadsheet-table`, `set-spreadsheet-table`, `add-spreadsheet-auto-filter`, -`set-spreadsheet-auto-filter`, `sort-spreadsheet-range`, `merge-cells`, -`unmerge-cells`, +`set-spreadsheet-auto-filter`, `sort-spreadsheet-range`, +`import-spreadsheet-delimited`, `set-spreadsheet-frozen-pane`, `merge-cells`, +`unmerge-cells`, `recalculate-spreadsheet-formulas`, `set-hyperlink`, `set-comment`, `set-table-column-width`, `set-cell-value`, `add-paragraph`, `add-table`, `add-table-row`, `add-table-column`, `add-table-cell`, @@ -1012,6 +1095,39 @@ follow the row permutation. The worksheet used dimension is recomputed and chart caches are cleared. Exact replay emits the same `sort-spreadsheet-range` mutation. +A Spreadsheet delimited import embeds bounded content in the typed batch or MCP +request; filesystem paths remain a CLI-only concern: + +```json +{ + "operation": "import-spreadsheet-delimited", + "sheet": "/Sheet1", + "import": { + "content": "Name,Amount,Date\nAlpha,42,2026-07-17", + "format": "csv", + "header": true, + "startCell": "A1" + } +} +``` + +`format` is `csv` or `tsv`. Header mode replaces the worksheet AutoFilter range +and its canonical frozen pane in the same transaction, so inspect those nodes +before importing into a populated sheet. Set a pane independently with +`set-spreadsheet-frozen-pane` and remove it through `/Sheet/freeze`: + +```json +{ + "operation": "set-spreadsheet-frozen-pane", + "sheet": "/Sheet1", + "pane": { + "frozenRows": 1, + "frozenColumns": 0, + "topLeftCell": "A2" + } +} +``` + A Spreadsheet table mutation uses one complete ListObject value. CLI `set` preserves omitted fields; batch, Rust, and standard MCP replacements supply the complete table: @@ -1107,8 +1223,9 @@ that current typed mutations can reproduce byte-for-byte at the OOXML part-map level: plain Word paragraphs and rectangular tables, Spreadsheet worksheets, typed defined names, typed cells, typed worksheet/table AutoFilters, typed ListObject tables, stable physical row order with supported typed sort state, -merged ranges, typed data-validation rules, and canonical typed -conditional-format rules without cached formula results; +canonical frozen panes and import date styles, merged ranges, typed +data-validation rules, canonical typed conditional-format rules, and natively +recalculable formula caches and canonical cached dynamic-array spills; plus Presentation slides with plain one-run text shapes and canonical basic tables. Headers, notes, @@ -1391,10 +1508,26 @@ Typed Spreadsheet content values use an explicit nested type, for example: } ``` -Formula mutation stores OOXML formula text and marks the workbook for a full -recalculation when opened. Structural edits rewrite supported A1 references -without evaluating formulas. A complete formula parser, dependency graph, and -evaluator remain a separate rich-Spreadsheet delivery gate. +Formula mutation removes one optional leading `=`, parses the bounded body into +a source-spanned typed AST, stores the original normalized OOXML formula text, +and marks the workbook for recalculation. It does not implicitly calculate the +workbook. The parser covers scalar and error literals, Excel operator +precedence, function calls and omitted arguments, parentheses and array +constants, names and structured references, A1 cell/row/column references with +quoted, 3D, or external qualifiers, and range/intersection/union operators. +Invalid syntax returns `use.office.spreadsheet_formula_invalid` with zero-based +UTF-8 byte and character offsets before any package mutation. Structural edits +continue to rewrite supported A1 references without evaluating formulas. + +`NativeOfficeDocument::formula_dependency_graph` and +`calculate_spreadsheet_formulas` provide read-only graph and calculation +results. `NativeOfficeEditor::recalculate_spreadsheet_formulas`, the +`recalculate-spreadsheet-formulas` batch/MCP mutation, exact replay, and the +CLI command above atomically persist supported cached values and spills. Excel +formula breadth beyond the closed native registry, disjoint or non-canonical +structured-reference forms, qualified functions, external-workbook +calculation, and full cross-application conformance remain rich-Spreadsheet +delivery gates. The native package, semantic, and editor APIs are available directly to Rust callers: diff --git a/crates/office/skills/a3s-use-office/SKILL.md b/crates/office/skills/a3s-use-office/SKILL.md index 6fe1056f..07de1b29 100644 --- a/crates/office/skills/a3s-use-office/SKILL.md +++ b/crates/office/skills/a3s-use-office/SKILL.md @@ -71,8 +71,16 @@ available. is unavoidable, inspect the exact part, preserve its root QName, write to a distinct output, and validate the result. - Do not evaluate formulas through a shell or general-purpose script runtime. - Native formula writes request spreadsheet recalculation; they do not promise - a computed cached value. + Native cell-formula writes validate and store the expression but do not + compute a cached value implicitly. When fresh results are required, use + `office native recalculate` or the typed + `recalculate-spreadsheet-formulas` mutation. The closed native function + registry must reject unsupported functions instead of falling back to code + execution. +- Treat dynamic-array spill children as read-only calculated output. Find and + edit or remove the formula anchor whose `formulaRef` contains the child; + recalculation, cache writes, spill cleanup, and every sibling mutation in + the batch roll back together on failure. - Treat external OOXML relationships as inert. Do not fetch linked resources while inspecting or rendering a document. Native hyperlink writes accept only absolute HTTP, HTTPS, or mailto URIs without embedded credentials. @@ -101,8 +109,18 @@ available. `namedrange` first and use the returned `@name` plus `@scope` path for update/remove. Do not edit `_xlnm.*` or `Slicer_*` names, collide with a table name, add a formula-bar leading `=`, or use raw XML to bypass a typed - identity/ref error. Defined-name formulas are stored and marked for - recalculation, not evaluated by A3S. + identity/ref error. Defined-name mutations store the definition and request + recalculation; supported names referenced by cell formulas are resolved only + during an explicit native recalculation pass. +- Import CSV or TSV only through the bounded typed import. Use exactly one + regular source file or `--stdin`, make `--format` explicit for stdin, and + inspect the target worksheet, `/Sheet/autofilter`, and `/Sheet/freeze` before + enabling `--header`: header mode intentionally replaces the worksheet filter + range and canonical frozen pane in one transaction. Explicit empty fields + clear existing cells; missing trailing fields in ragged rows do not. Treat + inferred formulas as parsed but not implicitly recalculated; run the explicit + native pass when cached values are required. Never bypass a malformed-quote, + formula-syntax, range, type, or unknown-view error with `raw-set`. - Treat Spreadsheet AutoFilters as typed worksheet or table structure. Query `autofilter` or `filtercolumn` first and inspect `nativeMutable`; use the stable `/Sheet/autofilter` or `/Sheet/table[N]` path for updates. Every @@ -126,9 +144,14 @@ available. not overlap another table, a merge, or a worksheet AutoFilter, and do not use raw XML to bypass `nativeMutable=false` or an unknown-content/relationship error. Table criteria use the same typed filter-column values as worksheet - AutoFilters. Exact mutable table or data ranges without totals rows can use - the separate physical sort contract; unsupported embedded/imported sort - state remains non-mutable. + AutoFilters. Table set automatically rewrites common explicit structured + references and provably owned table-local column references when aliases or + position-mapped columns change. Do not bypass + `use.office.spreadsheet_table_formula_rewrite_unsupported` for unsafe local + geometry or ownership, or `use.office.spreadsheet_table_referenced` when + removal is blocked. Exact mutable table or data ranges without totals rows + can use the separate physical sort contract; unsupported embedded/imported + sort state remains non-mutable. - Keep the default OfficeCLI compatibility route separate from the native engine. Do not depend on OfficeCLI's private resident protocol. @@ -152,6 +175,12 @@ typed thresholds, stable paths, semantic queries, and exact canonical replay. It owns workbook-global and worksheet-local Spreadsheet defined names with stable scoped paths, typed add/set/remove, semantic readback, and exact replay. It owns typed Spreadsheet +formula parsing, bounded dependency graphs, a closed typed function registry, +read-only calculation, atomic cached-value and dynamic-array spill writeback, +CLI/MCP/batch recalculation, and exact replay. It owns typed Spreadsheet +CSV/TSV import with bounded strict parsing, typed cell inference, explicit +empty-cell semantics, optional header AutoFilter/frozen-pane setup, semantic +`/Sheet/freeze` state, and exact canonical replay. It owns typed Spreadsheet worksheet and table AutoFilters with closed value, comparison, top/bottom, and dynamic criteria, stable filter paths, add/set/remove, and exact replay. It owns stable ordered multi-key Spreadsheet physical sorting over an explicit or @@ -168,12 +197,12 @@ cells or bounded ranges and internal locations, and external Presentation shape clicks or internal jumps to existing slides. Remaining boundaries include modern threaded comments, replies/resolution, writable comment dates, rich comment bodies, Word header/footer comment anchors, -gradient/pattern/theme fills, advanced x14 conditional-format visuals, named styles, complete formula -calculation, formula-bearing or table-totals sorting, table calculated +gradient/pattern/theme fills, advanced x14 conditional-format visuals, named +styles, complete Excel function/structured-reference/external-workbook formula +compatibility, formula-bearing or table-totals sorting, table calculated columns/totals functions, date-group/color/icon filters and unsupported embedded/imported sort-state variants, custom table styles, query -tables/external data, advanced charts, pivots, -and media, +tables/external data, advanced charts, pivots, and media, interactive preview editing/annotations, and full Office layout fidelity. Fail closed or use the explicit compatibility route rather than inventing unsupported native behavior. diff --git a/crates/office/skills/a3s-use-office/references/mcp.md b/crates/office/skills/a3s-use-office/references/mcp.md index 5d1da61b..3c2102c8 100644 --- a/crates/office/skills/a3s-use-office/references/mcp.md +++ b/crates/office/skills/a3s-use-office/references/mcp.md @@ -117,6 +117,50 @@ mutations. Unknown fields, invalid values, empty format objects, and non-Spreadsheet targets fail the entire in-memory batch; no change persists until `office_save`. +Spreadsheet cell formulas use `set-cell-value` with +`{"type":"formula","expression":"SUM(A1:B2)"}`. One optional leading `=` is +removed before storage. The bounded native parser validates literals, +operators, calls, names, structured references, qualified A1 references, and +range/intersection/union syntax. Invalid syntax returns +`use.office.spreadsheet_formula_invalid` with zero-based `byteOffset` and +`characterOffset` details and rolls back every mutation in that +`office_apply_batch`. Successful writes invalidate stale calculation caches and +request recalculation; they do not calculate implicitly. + +Add the explicit recalculation mutation after dependent formula writes when +fresh cached values are required: + +```json +{ + "session": "workbook", + "mutations": [ + { + "operation": "set-cell-value", + "path": "/Sheet1/C1", + "value": {"type": "formula", "expression": "SEQUENCE(2,2,1,1)"} + }, + { + "operation": "recalculate-spreadsheet-formulas" + } + ] +} +``` + +The result includes one `spreadsheetCalculations` receipt with +`formulaCount`, `spillCellCount`, deterministic `calculationOrder`, and typed +calculated cells. Calculation and OOXML cache/spill writeback are part of the +same atomic in-memory batch; a cycle, unsupported function, qualified function, +missing structured-reference table, column, or requested header/totals row, +external-workbook reference, or blocked storage condition rolls back every +sibling mutation. `Table[Column]`, contiguous table column ranges, common +`#All`/`#Data`/`#Headers`/`#Totals` row items, current-row `@` forms, and +table-local current-row references from inside a table are supported. +Spreadsheet error results remain typed cell values. Dynamic-array spill +children are read-only; mutate their formula anchor. The closed native registry +and limits are documented in +[spreadsheet.md](spreadsheet.md#values-and-formulas); the server never invokes +a shell, script runtime, or external workbook. + Spreadsheet merged cells use the separate `merge-cells` and `unmerge-cells` mutations: @@ -179,7 +223,8 @@ operator, and only `between` or `notBetween` accept and require `formula2`. Rules and ranges are bounded, normalized, and globally non-overlapping within one worksheet. Invalid formulas, flags, messages, XML text, ranges, or overlap fail the complete `office_apply_batch`. Inline lists, ISO dates, and clock -times are normalized but formulas are never evaluated. Query +times are normalized, but data-validation formula predicates are not executed +by the validation feature or the cell-formula recalculation pass. Query `dataValidation[type=list]` or call `office_get` on the returned path for unsaved semantic readback. Covered observed and virtual blank cells expose `dataValidation` and `validationType`. Updates retain unknown attributes and @@ -273,11 +318,12 @@ identity, ListObject table-name collisions, reserved `_xlnm.*`/`Slicer_*` names, and unsupported cross-workbook refs are validated before mutation. Workbook-scoped bare A1 refs are rejected; worksheet-local bare A1 refs are qualified automatically by the domain layer. The mutation requests workbook -recalculation but does not evaluate the expression. Unknown OOXML attributes -are retained, while unknown content that cannot be preserved fails the whole -batch. Call `office_get` or `office_query` before `office_save` to verify the -unsaved scoped value, then save explicitly. Closing a dirty session still -requires save or explicit discard. +recalculation but does not calculate by itself; a later explicit native +recalculation resolves supported names referenced by cell formulas. Unknown +OOXML attributes are retained, while unknown content that cannot be preserved +fails the whole batch. Call `office_get` or `office_query` before `office_save` +to verify the unsaved scoped value, then save explicitly. Closing a dirty +session still requires save or explicit discard. Spreadsheet worksheet AutoFilters use the separate `add-spreadsheet-auto-filter` and `set-spreadsheet-auto-filter` mutations. @@ -323,6 +369,60 @@ color/icon, extension, unknown-content, and embedded sort-state imports fail closed. Physical sorting is the separate mutation below and does not flatten an unsupported imported AutoFilter. +Spreadsheet delimited import embeds bounded UTF-8 content directly in the +typed mutation; filesystem paths remain at the CLI boundary: + +```json +{ + "session": "workbook", + "mutations": [{ + "operation": "import-spreadsheet-delimited", + "sheet": "/Sheet1", + "import": { + "content": "Name,Amount,Date\nAlpha,42,2026-07-17", + "format": "csv", + "header": true, + "startCell": "A1" + } + }] +} +``` + +`format` is `csv` or `tsv`; omitted `startCell` defaults to `A1`. Input is +limited to 8 MiB and a 100,000-cell rectangular extent. Malformed quoting, +invalid geometry or typed values, and unsupported target state fail the whole +in-memory batch. Explicit empty fields clear existing target values; ragged +missing trailing fields preserve them. Formula, finite-number, boolean, and ISO +date/time inference is deterministic, but import does not implicitly calculate +formulas. Append `recalculate-spreadsheet-formulas` to the same batch when +fresh caches are required. + +When `header=true`, the import atomically adds or replaces the worksheet +AutoFilter and canonical frozen pane. Inspect `/Sheet1/autofilter` and +`/Sheet1/freeze` before using header mode on a populated worksheet. A pane can +also be set independently: + +```json +{ + "session": "workbook", + "mutations": [{ + "operation": "set-spreadsheet-frozen-pane", + "sheet": "/Sheet1", + "pane": { + "frozenRows": 1, + "frozenColumns": 0, + "topLeftCell": "A2" + } + }] +} +``` + +Use `office_get` or `office_query` for unsaved readback and ordinary typed +`remove` on `/Sheet1/freeze` for deletion. Do not mutate a pane whose semantic +`nativeMutable` value is false. See +[spreadsheet.md](spreadsheet.md#delimited-import-and-frozen-panes) for parsing, +typing, and preservation boundaries. + Spreadsheet physical sorting uses `sort-spreadsheet-range` inside the same atomic `office_apply_batch` boundary: @@ -408,6 +508,15 @@ Names, columns, built-in style families/numbers, flags, table/defined-name identity, table/merge/worksheet-AutoFilter overlap, and relationship ownership are validated before mutation. +When table aliases or position-mapped columns change, `set-spreadsheet-table` +atomically rewrites common structured references in cells, defined names, +conditional formats, data validations, charts, and table-formula carriers. +String literals and external-workbook references are preserved. Unsafe +table-local rewrites or geometry changes fail with +`use.office.spreadsheet_table_formula_rewrite_unsupported`; removing a table +with a remaining structured reference fails with +`use.office.spreadsheet_table_referenced`. + Use `office_get` with depth 1 or `office_query` with `table[name=Sales]` to inspect the unsaved table and its column children. Do not replace a node whose semantic `nativeMutable` flag is false. Header stamping, OPC table parts, and diff --git a/crates/office/skills/a3s-use-office/references/spreadsheet.md b/crates/office/skills/a3s-use-office/references/spreadsheet.md index 672d86da..3b75be3c 100644 --- a/crates/office/skills/a3s-use-office/references/spreadsheet.md +++ b/crates/office/skills/a3s-use-office/references/spreadsheet.md @@ -7,6 +7,7 @@ Use stable worksheet and A1 paths such as `/Sheet1`, `/Sheet1/A1`, and - [Inspect](#inspect) - [Values and Formulas](#values-and-formulas) +- [Delimited Import and Frozen Panes](#delimited-import-and-frozen-panes) - [Cell Text Formatting](#cell-text-formatting) - [Cell Presentation Formatting](#cell-presentation-formatting) - [Merged Cells](#merged-cells) @@ -37,6 +38,7 @@ a3s use office native set workbook.xlsx /Sheet1/A1:C20 --find Draft --replace Fi a3s use office native set workbook.xlsx /Sheet1/B1 --number 42.5 --json a3s use office native set workbook.xlsx /Sheet1/C1 --boolean true --json a3s use office native set workbook.xlsx /Sheet1/D1 --formula 'SUM(B1:B12)' --json +a3s use office native recalculate workbook.xlsx --output calculated.xlsx --json a3s use office native set workbook.xlsx /Sheet1/E1 --url https://example.com/data --display Data --tooltip 'Open data' --json a3s use office native set workbook.xlsx /Sheet1/F1 --location 'Sheet1!B2' --display B2 --json a3s use office native set workbook.xlsx /Sheet1/G2:H4 --url https://example.com/range --display Range --json @@ -56,10 +58,54 @@ the scope. Rich runs and unknown XML survive, and phonetic text is excluded. Numeric, boolean, formula, and error values are not coerced. Zero matches are reported as an unchanged success. -Formula writes store validated formula text, invalidate stale calculation -caches, and request application recalculation. The native engine does not yet -provide a complete formula evaluator. Check `formula_not_evaluated` and -`formula_eval_error` issue records before delivery. +Formula writes remove one optional leading `=`, parse the body with bounded +Excel operator/reference syntax, store the normalized formula text, invalidate +stale calculation caches, and request recalculation. They do not calculate +implicitly. A syntax failure returns +`use.office.spreadsheet_formula_invalid` with byte and character offsets and +leaves the document unchanged. + +Run `office native recalculate` in place or with `--output` to build the +dependency graph, calculate supported formulas, and atomically write typed +cached values and dynamic-array spills. The same operation is available as the +`recalculate-spreadsheet-formulas` batch/MCP mutation and as read-only or +writeback Rust APIs. Supported functions are `SUM`, `AVERAGE`, `MIN`, `MAX`, +`COUNT`, `COUNTA`, `ABS`, `SQRT`, `POWER`, `MOD`, `ROUND`, `IF`, `IFERROR`, +`AND`, `OR`, `NOT`, `CONCAT`, `CONCATENATE`, `ROW`, `COLUMN`, `SEQUENCE`, +`TRANSPOSE`, `PI`, and `NA`. Cross-sheet ranges, scoped names, typed errors, +array broadcasting, spill references, and ordinary Excel operators are +supported. ListObject structured references resolve a table `name` or +`displayName`: `Sales[Qty]` selects one data column, +`Sales[[Qty]:[Price]]` selects a contiguous data-column range, and `#All`, +`#Data`, `#Headers`, or `#Totals` selects structural rows. `Sales[@Qty]`, +`Sales[[#This Row],[Qty]]`, and table-local `[@Qty]` select the current data +row; table-local forms require the formula cell to be inside the inferred +table. + +Spill children are read-only; update or remove the anchor instead. A blocked +spill produces typed `#SPILL!`, while formula error values such as `#DIV/0!` +remain typed cell results. Circular dependencies, unsupported or qualified +functions, missing tables, columns, or requested header/totals rows, disjoint +or non-canonical structured-reference forms, and external-workbook reads fail +with stable errors and leave the complete mutation batch unchanged. No shell, +script runtime, or external workbook is invoked. Limits are 8,192 formula +characters, depth 128 across both AST and nested named-reference resolution, +8,192 AST nodes, 100,000 reference areas per value, 100,000 graph formulas, +1,000,000 dependency edges, 1,000,000 graph reference visits, 100,000 +materialized cells per array or function call, 100,000 cumulative spill +children per pass, 200,000 OOXML cell writes, and 1 MiB per text result. All +formula text results together are limited to 8 MiB per pass. Check +`formula_not_evaluated` and `formula_eval_error` issue records after the pass. + +Semantic cell reads expose string-valued `formulaCached` on formula anchors and +`valuePresent` on every cell. A recalculated anchor reports +`formulaCached=true`; a formula stored without `` reports `false`. Spill +children contain cached values but no independent `formula` field. + +Exact replay accepts canonical formula storage and canonical array anchors only +when the array result is natively cached. It fails closed with +`use.office.dump_unsupported` for non-reproducible physical storage such as +explicit `t="normal"` formulas and uncached or malformed array anchors. Hyperlinks target one cell or a bounded rectangular range. A missing single cell is auto-created; a range link neither creates cells nor rewrites their @@ -79,6 +125,77 @@ and slide coordinates instead of ignoring them. Native removal also cleans up the matching VML note shape and removes unused comment/VML parts. Threaded comments, replies, writable dates, and rich bodies are not yet native. +## Delimited Import and Frozen Panes + +Import one bounded UTF-8 CSV or TSV source into an existing worksheet: + +```bash +# .tsv and .tab infer TSV; every other file extension defaults to CSV. +a3s use office native import workbook.xlsx /Sheet1 source.csv \ + --header \ + --start-cell B2 \ + --json + +# Stdin is bounded too. State the format instead of relying on its CSV default. +a3s use office native import workbook.xlsx /Sheet1 \ + --stdin \ + --format tsv \ + --output imported.xlsx \ + --json +``` + +Supply exactly one positional source, `--file `, or `--stdin`. Files +must be regular, non-symlink files. One request accepts at most 8 MiB and a +100,000-cell rectangular target within Excel's row and column bounds. The +parser accepts a leading UTF-8 BOM, CRLF, quoted delimiters, embedded quoted +newlines, and doubled quotes. It rejects unclosed quotes, quotes inside +unquoted fields, and non-boundary content after a closing quote rather than +guessing. + +An explicit empty field clears an existing target cell value while retaining +its unrelated style and extension content. A missing trailing field in a +ragged source row leaves that target cell unchanged, and a blank target is not +materialized just to represent emptiness. Import infers leading-`=` formulas, +finite numbers, booleans, ISO dates/times, and otherwise text. Dates honor the +workbook's 1900/1904 date system and receive the canonical native date number +format. Inferred formulas pass the same bounded syntax parser as direct cell +writes, are stored, and are marked for recalculation. Import does not calculate +them implicitly; run `office native recalculate` when fresh cached values are +required. + +`--header` treats the first imported row as headers. In the same atomic +transaction it adds or replaces the worksheet AutoFilter over the imported +extent and adds or replaces one canonical frozen pane below the header. Inspect +existing `/Sheet1/autofilter` and `/Sheet1/freeze` state first when importing +into a populated worksheet. + +Read or remove the frozen pane through its stable semantic path: + +```bash +a3s use office native get workbook.xlsx /Sheet1/freeze --json +a3s use office native query workbook.xlsx frozen-pane --json +a3s use office native remove workbook.xlsx /Sheet1/freeze --json +``` + +Rust, versioned batch, and standard MCP can set a canonical pane independently: + +```json +{ + "operation": "set-spreadsheet-frozen-pane", + "sheet": "/Sheet1", + "pane": { + "frozenRows": 1, + "frozenColumns": 0, + "topLeftCell": "A2" + } +} +``` + +`topLeftCell` must be below and to the right of every frozen split. Imported +split panes, vendor attributes, unknown children, or unsupported view state +remain readable with `nativeMutable=false` and fail closed on set/remove. +Strict and transitional SpreadsheetML are preserved. + ## Cell Text Formatting ```bash @@ -391,6 +508,18 @@ a3s use office native set workbook.xlsx '/Sheet1/table[1]' \ a3s use office native remove workbook.xlsx '/Sheet1/table[1]' --json ``` +Table `set` rewrites common structured references when `name`, effective +`displayName`, or position-mapped column names change. The audit covers cell +formulas, workbook defined names, conditional-format and data-validation +formulas, charts, and formula carriers in table parts. String literals and +external-workbook references remain unchanged. Table-local forms such as +`[@Qty]` are rewritten only with provable ListObject ownership; an unsafe local +rewrite or local reference across a range/header/totals-row change fails with +`use.office.spreadsheet_table_formula_rewrite_unsupported`. Removing a table +still targeted by a structured reference fails with +`use.office.spreadsheet_table_referenced`. Both failures roll back the complete +mutation. + Provide exactly one non-empty, case-insensitively unique column name for every range column. Table `name` and optional `displayName` use Excel identifier grammar, are limited to 255 characters, may not resemble A1/R1C1 references, @@ -512,7 +641,9 @@ rejected; use cells as the list source instead. Date inputs in valid workbook's declared 1900 or 1904 date system. Time inputs in `HH:MM` or `HH:MM:SS` form become day fractions. Range, defined-name, dynamic spill, and function sources such as `INDIRECT(...)` remain formulas. Other formula text is -stored after removing one optional leading `=` and is never evaluated by A3S. +stored after removing one optional leading `=`; data-validation rule predicates +are not executed by either the validation writer or the cell-formula +recalculation pass. Each rule accepts 1–1,024 normalized rectangular A1 areas and a worksheet accepts at most 65,534 rules. Formula fields are limited to 255 characters; @@ -556,8 +687,8 @@ children would be lost, and final removal fails if unknown collection data would be discarded. Strict/transitional OOXML, atomic batch rollback, and exact replay are supported. This capability does not add table calculated columns/totals functions, date-group/color/icon filters, unsupported imported -sort-state variants, charts, pivots, formula evaluation, or Excel layout -fidelity. +sort-state variants, charts, pivots, data-validation predicate execution, or +Excel layout fidelity. ## Conditional Formatting @@ -668,9 +799,10 @@ survive a set/remove operation fails closed. Imported multi-rule carriers share one range: keep that range unchanged when updating one child rule. Canonical replay, atomic rollback, CLI, and standard MCP are supported. This -does not calculate formulas, reproduce Excel's rendered appearance, or support -x14-only negative data-bar axes/colors, custom icon sets, table/chart/pivot -formatting, or complete OfficeCLI/Spreadsheet parity. +conditional-format feature does not evaluate rule formulas or reproduce +Excel's rendered appearance, and it does not support x14-only negative data-bar +axes/colors, custom icon sets, table/chart/pivot formatting, or complete +OfficeCLI/Spreadsheet parity. ## Named Ranges @@ -733,15 +865,16 @@ The identity is case-insensitively unique by `(name, scope)`. A defined name also may not collide with a ListObject table `name` or `displayName`. Do not edit or remove `_xlnm.*` print/filter definitions or `Slicer_*` sentinels; manage the owning typed feature instead. `--volatile true` maps to the OOXML -defined-name function flag and requests recalculation. No named-range formula -is evaluated by A3S. +defined-name function flag and requests recalculation. The name mutation does +not itself calculate anything; supported names referenced by cell formulas are +resolved by an explicit native recalculation pass. Batch, standard MCP, and Rust use one complete typed value for add/set and ordinary typed `remove` for deletion. The writer preserves strict/transitional SpreadsheetML and unknown attributes. Unknown collection or child content fails closed when an edit cannot retain it. Exact replay includes supported defined names. This remains defined-name lifecycle support, not external-link -authoring, formula evaluation, or complete Spreadsheet parity. +authoring or complete Spreadsheet parity. ## Structure @@ -757,8 +890,8 @@ a3s use office native add workbook.xlsx /Sheet1/A1 --type picture --input chart. Supported structural edits rewrite bounded A1 references and related metadata. Pivot-table changes, unsafe 3D references, x14-only conditional-format -extensions, full chart authoring, and complete recalculation remain outside the -native subset and fail closed where safety cannot be proven. +extensions, full chart authoring, and complete Excel formula compatibility +remain outside the native subset and fail closed where safety cannot be proven. ## Verify @@ -772,4 +905,4 @@ a3s use office native watch workbook.xlsx --port 0 HTML, SVG, and screenshots are sparse semantic previews, not Excel layout or print fidelity. Watch reloads saved revisions; it does not provide inline cell -editing or calculate formulas. +editing or trigger formula recalculation. diff --git a/crates/office/src/editor.rs b/crates/office/src/editor.rs index c023f96b..3f8b0110 100644 --- a/crates/office/src/editor.rs +++ b/crates/office/src/editor.rs @@ -47,14 +47,17 @@ pub use types::{ NativeSpreadsheetConditionalFormatThresholdKind, NativeSpreadsheetConditionalFormatTimePeriod, NativeSpreadsheetDataValidation, NativeSpreadsheetDataValidationErrorStyle, NativeSpreadsheetDataValidationOperator, NativeSpreadsheetDataValidationType, + NativeSpreadsheetDelimitedFormat, NativeSpreadsheetDelimitedImport, NativeSpreadsheetDifferentialFormat, NativeSpreadsheetDynamicFilter, NativeSpreadsheetFill, - NativeSpreadsheetFilterColumn, NativeSpreadsheetFilterCriteria, NativeSpreadsheetNamedRange, - NativeSpreadsheetNamedRangeScope, NativeSpreadsheetReadingOrder, NativeSpreadsheetSort, - NativeSpreadsheetSortDirection, NativeSpreadsheetSortKey, NativeSpreadsheetTable, - NativeSpreadsheetTableColumn, NativeSpreadsheetTableStyle, NativeSpreadsheetVerticalAlignment, - SpreadsheetCellValue, MAX_NATIVE_OFFICE_FIND_BYTES, MAX_NATIVE_OFFICE_REPLACEMENT_BYTES, + NativeSpreadsheetFilterColumn, NativeSpreadsheetFilterCriteria, NativeSpreadsheetFrozenPane, + NativeSpreadsheetImportResult, NativeSpreadsheetNamedRange, NativeSpreadsheetNamedRangeScope, + NativeSpreadsheetReadingOrder, NativeSpreadsheetSort, NativeSpreadsheetSortDirection, + NativeSpreadsheetSortKey, NativeSpreadsheetTable, NativeSpreadsheetTableColumn, + NativeSpreadsheetTableStyle, NativeSpreadsheetVerticalAlignment, SpreadsheetCellValue, + MAX_NATIVE_OFFICE_FIND_BYTES, MAX_NATIVE_OFFICE_REPLACEMENT_BYTES, MAX_NATIVE_OFFICE_TEXT_MATCHES, MAX_NATIVE_OFFICE_TEXT_REPLACEMENT_OUTPUT_BYTES, - MAX_NATIVE_OFFICE_TEXT_SCOPE_CELLS, + MAX_NATIVE_OFFICE_TEXT_SCOPE_CELLS, MAX_NATIVE_SPREADSHEET_IMPORT_BYTES, + MAX_NATIVE_SPREADSHEET_IMPORT_CELLS, }; /// Loss-preserving OOXML editor with transactional in-memory batches. @@ -202,6 +205,24 @@ impl NativeOfficeEditor { Ok(()) } + /// Calculates every supported Spreadsheet formula and atomically writes + /// typed cached values and dynamic-array spill cells into the package. + pub fn recalculate_spreadsheet_formulas( + &mut self, + ) -> UseResult { + let result = self.apply_batch(&[NativeOfficeMutation::RecalculateSpreadsheetFormulas])?; + result + .spreadsheet_calculations + .into_iter() + .next() + .ok_or_else(|| { + editor_error( + "use.office.batch_validation_failed", + "Native Spreadsheet recalculation returned no calculation receipt.", + ) + }) + } + /// Adds one complete typed Spreadsheet ListObject table. pub fn add_spreadsheet_table( &mut self, @@ -262,6 +283,40 @@ impl NativeOfficeEditor { }) } + /// Imports bounded CSV or TSV content into one Spreadsheet worksheet. + pub fn import_spreadsheet_delimited( + &mut self, + sheet: impl Into, + import: NativeSpreadsheetDelimitedImport, + ) -> UseResult { + let result = self.apply_batch(&[NativeOfficeMutation::ImportSpreadsheetDelimited { + sheet: sheet.into(), + import, + }])?; + result + .spreadsheet_imports + .into_iter() + .next() + .ok_or_else(|| { + editor_error( + "use.office.batch_validation_failed", + "Native Spreadsheet import returned no receipt.", + ) + }) + } + + /// Creates or replaces one canonical frozen pane on a Spreadsheet sheet. + pub fn set_spreadsheet_frozen_pane( + &mut self, + sheet: impl Into, + pane: NativeSpreadsheetFrozenPane, + ) -> UseResult { + self.single_path(NativeOfficeMutation::SetSpreadsheetFrozenPane { + sheet: sheet.into(), + pane, + }) + } + /// Adds one complete typed Spreadsheet defined name. pub fn add_named_range( &mut self, @@ -710,6 +765,8 @@ impl NativeOfficeEditor { created_parts: Vec::new(), created_images: Vec::new(), text_replacements: Vec::new(), + spreadsheet_imports: Vec::new(), + spreadsheet_calculations: Vec::new(), }); } let original = self.package.clone(); @@ -718,11 +775,15 @@ impl NativeOfficeEditor { let mut created_parts = Vec::new(); let mut created_images = Vec::new(); let mut text_replacements = Vec::new(); + let mut spreadsheet_imports = Vec::new(); + let mut spreadsheet_calculations = Vec::new(); for mutation in mutations { let mut created_part = None; let mut created_image = None; let mut swap = None; let mut text_replacement = None; + let mut spreadsheet_import = None; + let mut spreadsheet_calculation = None; let result = match mutation { NativeOfficeMutation::ReplaceText { path, replacement } => { text_replace::replace(&mut self.package, path, replacement).map(|receipt| { @@ -766,6 +827,12 @@ impl NativeOfficeEditor { spreadsheet::set_cell_value(&mut self.package, path, value) .map(|()| path.clone()) } + NativeOfficeMutation::RecalculateSpreadsheetFormulas => { + spreadsheet::recalculate_formulas(&mut self.package).map(|receipt| { + spreadsheet_calculation = Some(receipt); + "/".to_string() + }) + } NativeOfficeMutation::AddSpreadsheetTable { sheet, table } => { spreadsheet::add_table(&mut self.package, sheet, table) } @@ -781,6 +848,16 @@ impl NativeOfficeEditor { NativeOfficeMutation::SortSpreadsheetRange { path, sort } => { spreadsheet::sort_range(&mut self.package, path, sort) } + NativeOfficeMutation::ImportSpreadsheetDelimited { sheet, import } => { + spreadsheet::import_delimited(&mut self.package, sheet, import).map(|receipt| { + let path = receipt.path.clone(); + spreadsheet_import = Some(receipt); + path + }) + } + NativeOfficeMutation::SetSpreadsheetFrozenPane { sheet, pane } => { + spreadsheet::set_frozen_pane(&mut self.package, sheet, pane) + } NativeOfficeMutation::AddNamedRange { named_range } => { spreadsheet::add_named_range(&mut self.package, named_range) } @@ -953,6 +1030,12 @@ impl NativeOfficeEditor { if let Some(receipt) = text_replacement { text_replacements.push(receipt); } + if let Some(receipt) = spreadsheet_import { + spreadsheet_imports.push(receipt); + } + if let Some(receipt) = spreadsheet_calculation { + spreadsheet_calculations.push(receipt); + } } Err(error) => { self.package = original; @@ -974,6 +1057,8 @@ impl NativeOfficeEditor { created_parts, created_images, text_replacements, + spreadsheet_imports, + spreadsheet_calculations, }) } diff --git a/crates/office/src/editor/spreadsheet.rs b/crates/office/src/editor/spreadsheet.rs index dfc98664..2730cdff 100644 --- a/crates/office/src/editor/spreadsheet.rs +++ b/crates/office/src/editor/spreadsheet.rs @@ -18,18 +18,27 @@ mod auto_filter; mod conditional_formatting; mod data_validation; mod filter_xml; +mod formula; +mod import; mod merge; mod named_range; mod sort; mod structure; mod style; mod table; +mod view; mod worksheet; pub(super) use arrange::{copy_node, move_node, swap_nodes}; pub(super) use structure::{delete_columns, delete_rows, insert_columns, insert_rows}; pub(super) use worksheet::{copy_worksheet, move_worksheet, rename_worksheet}; +pub(super) fn recalculate_formulas( + package: &mut NativeOfficePackage, +) -> UseResult { + formula::recalculate(package) +} + pub(super) fn add_auto_filter( package: &mut NativeOfficePackage, sheet: &str, @@ -54,6 +63,22 @@ pub(super) fn sort_range( sort::sort(package, path, value) } +pub(super) fn set_frozen_pane( + package: &mut NativeOfficePackage, + sheet: &str, + pane: &super::NativeSpreadsheetFrozenPane, +) -> UseResult { + view::set(package, sheet, pane) +} + +pub(super) fn import_delimited( + package: &mut NativeOfficePackage, + sheet: &str, + import: &super::NativeSpreadsheetDelimitedImport, +) -> UseResult { + import::apply(package, sheet, import) +} + pub(super) fn add_conditional_format( package: &mut NativeOfficePackage, sheet: &str, @@ -191,6 +216,12 @@ pub(super) fn set_cell_value( })?; let part = package.xml_part(part_name)?; let index = index_xml(&part)?; + let sheet_data = index + .descendant("sheetData") + .ok_or_else(|| node_not_found(path))?; + let prepared = formula::prepare_for_value_write(&part, sheet_data, sheet, range)?; + let part = crate::LosslessXmlPart::parse(part_name.to_string(), prepared)?; + let index = index_xml(&part)?; let sheet_data = index .descendant("sheetData") .ok_or_else(|| node_not_found(path))?; @@ -201,7 +232,9 @@ pub(super) fn set_cell_value( } pub(super) fn remove(package: &mut NativeOfficePackage, path: &str) -> UseResult<()> { - if sort::is_path(path) { + if view::is_path(path) { + view::remove(package, path) + } else if sort::is_path(path) { sort::remove(package, path) } else if auto_filter::is_path(path) { auto_filter::remove(package, path) @@ -221,8 +254,6 @@ pub(super) fn remove(package: &mut NativeOfficePackage, path: &str) -> UseResult } fn remove_cell(package: &mut NativeOfficePackage, path: &str) -> UseResult<()> { - super::comment::remove_spreadsheet_range_comments(package, path)?; - super::hyperlink::remove_spreadsheet_range_links(package, path)?; let (sheet_path, reference) = path.rsplit_once('/').ok_or_else(|| node_not_found(path))?; let range = CellRange::parse(reference)?; validate_range_size(range)?; @@ -250,6 +281,15 @@ fn remove_cell(package: &mut NativeOfficePackage, path: &str) -> UseResult<()> { })?; let part = package.xml_part(part_name)?; let index = index_xml(&part)?; + let sheet_data = index + .descendant("sheetData") + .ok_or_else(|| node_not_found(path))?; + let prepared = formula::prepare_for_remove(&part, sheet_data, sheet, range)?; + package.set_part(part_name, prepared)?; + super::comment::remove_spreadsheet_range_comments(package, path)?; + super::hyperlink::remove_spreadsheet_range_links(package, path)?; + let part = package.xml_part(part_name)?; + let index = index_xml(&part)?; let sheet_data = index .descendant("sheetData") .ok_or_else(|| node_not_found(path))?; @@ -638,16 +678,8 @@ fn normalize_cell_value(value: &SpreadsheetCellValue) -> UseResult { - let expression = expression.strip_prefix('=').unwrap_or(expression); - if expression.is_empty() - || expression.chars().count() > 8_192 - || expression.chars().any(char::is_control) - { - return Err(editor_error( - "use.office.spreadsheet_formula_invalid", - "Spreadsheet formulas must contain 1-8192 non-control characters.", - )); - } + let expression = + crate::spreadsheet_formula::validate_and_normalize_formula(expression)?; Ok(SpreadsheetCellValue::Formula { expression: expression.to_string(), }) diff --git a/crates/office/src/editor/spreadsheet/formula.rs b/crates/office/src/editor/spreadsheet/formula.rs new file mode 100644 index 00000000..2b9f79f9 --- /dev/null +++ b/crates/office/src/editor/spreadsheet/formula.rs @@ -0,0 +1,169 @@ +mod planning; +mod write; + +use std::collections::BTreeMap; + +use a3s_use_core::{UseError, UseResult}; + +use crate::semantic::{DocumentNode, NativeOfficeDocument}; +use crate::spreadsheet_reference::{CellRange, CellReference}; +use crate::xml_edit::{index_xml, IndexedXmlElement}; +use crate::{ + NativeOfficePackage, SpreadsheetFormulaCalculation, SpreadsheetFormulaValue, + MAX_SPREADSHEET_FORMULA_CELLS, MAX_SPREADSHEET_FORMULA_SPILL_CELLS, +}; + +use super::{editor_error, update_dimension}; +use planning::{plan_writes, worksheet_cells}; +use write::{apply_cell_writes, mark_workbook_calculated}; + +const MAX_CALCULATION_WRITES: usize = + MAX_SPREADSHEET_FORMULA_CELLS + MAX_SPREADSHEET_FORMULA_SPILL_CELLS; + +#[derive(Debug, Clone)] +enum CellWrite { + Clear, + Cached(SpreadsheetFormulaValue), + Formula { + expression: String, + value: SpreadsheetFormulaValue, + spill_range: Option, + }, +} + +pub(super) fn prepare_for_value_write( + part: &crate::LosslessXmlPart, + sheet_data: &IndexedXmlElement, + sheet: &DocumentNode, + target: CellRange, +) -> UseResult> { + prepare_spill_edit(part, sheet_data, sheet, target, false) +} + +pub(super) fn prepare_for_remove( + part: &crate::LosslessXmlPart, + sheet_data: &IndexedXmlElement, + sheet: &DocumentNode, + target: CellRange, +) -> UseResult> { + prepare_spill_edit(part, sheet_data, sheet, target, true) +} + +fn prepare_spill_edit( + part: &crate::LosslessXmlPart, + sheet_data: &IndexedXmlElement, + sheet: &DocumentNode, + target: CellRange, + remove: bool, +) -> UseResult> { + let cells = worksheet_cells(sheet)?; + let mut writes = BTreeMap::new(); + for (anchor, cell) in &cells { + let Some(reference) = cell.format.get("formulaRef") else { + continue; + }; + let spill = CellRange::parse(reference)?; + if !target.intersects(spill) { + continue; + } + if !target.contains(*anchor) + || (!remove && !intersection_is_anchor_only(target, spill, *anchor)) + { + return Err(calculation_storage_error( + "use.office.spreadsheet_formula_spill_cell_read_only", + format!( + "Cell range '{}' intersects spill '{}' outside formula anchor '{}'.", + target.a1(), + spill.a1(), + anchor.a1() + ), + ) + .with_suggestion( + "Edit or remove the spill formula anchor; spilled result cells are read-only.", + )); + } + let spill_cells = spill.cell_count()?; + if spill_cells > MAX_SPREADSHEET_FORMULA_SPILL_CELLS { + return Err(calculation_write_limit().with_detail("cells", spill_cells)); + } + for row in spill.start.row..=spill.end.row { + for column in spill.start.column..=spill.end.column { + let reference = CellReference { column, row }; + if reference != *anchor { + insert_clear_write(&mut writes, reference)?; + } + } + } + } + if writes.is_empty() { + Ok(part.raw().to_vec()) + } else { + apply_cell_writes(part, sheet_data, &writes) + } +} + +fn intersection_is_anchor_only(left: CellRange, right: CellRange, anchor: CellReference) -> bool { + let start_column = left.start.column.max(right.start.column); + let start_row = left.start.row.max(right.start.row); + let end_column = left.end.column.min(right.end.column); + let end_row = left.end.row.min(right.end.row); + start_column == anchor.column + && end_column == anchor.column + && start_row == anchor.row + && end_row == anchor.row +} + +pub(super) fn recalculate( + package: &mut NativeOfficePackage, +) -> UseResult { + let document = NativeOfficeDocument::from_package(package.clone())?; + let calculation = document.calculate_spreadsheet_formulas()?; + let plans = plan_writes(&document, &calculation)?; + for (part_name, writes) in plans { + if writes.is_empty() { + continue; + } + let part = package.xml_part(&part_name)?; + let index = index_xml(&part)?; + let sheet_data = index.descendant("sheetData").ok_or_else(|| { + calculation_storage_error( + "use.office.spreadsheet_formula_storage_invalid", + format!("Worksheet part '{part_name}' has no sheetData element."), + ) + })?; + let edited = apply_cell_writes(&part, sheet_data, &writes)?; + let edited = update_dimension(&part_name, edited)?; + package.set_part(&part_name, edited)?; + } + mark_workbook_calculated(package)?; + Ok(calculation) +} + +fn calculation_write_limit() -> UseError { + calculation_storage_error( + "use.office.spreadsheet_formula_write_limit", + format!("Native formula recalculation writes at most {MAX_CALCULATION_WRITES} cells."), + ) +} + +fn insert_clear_write( + writes: &mut BTreeMap, + reference: CellReference, +) -> UseResult<()> { + if writes.contains_key(&reference) { + return Ok(()); + } + let cells = writes + .len() + .checked_add(1) + .ok_or_else(calculation_write_limit)?; + if cells > MAX_CALCULATION_WRITES { + return Err(calculation_write_limit().with_detail("cells", cells)); + } + writes.insert(reference, CellWrite::Clear); + Ok(()) +} + +fn calculation_storage_error(code: &str, message: impl Into) -> UseError { + editor_error(code, message) +} diff --git a/crates/office/src/editor/spreadsheet/formula/planning.rs b/crates/office/src/editor/spreadsheet/formula/planning.rs new file mode 100644 index 00000000..ea2f17f2 --- /dev/null +++ b/crates/office/src/editor/spreadsheet/formula/planning.rs @@ -0,0 +1,337 @@ +use std::collections::BTreeMap; + +use a3s_use_core::UseResult; + +use crate::semantic::{DocumentNode, NativeOfficeDocument, OfficeNodeType}; +use crate::spreadsheet_reference::{CellRange, CellReference}; +use crate::{ + SpreadsheetFormulaCalculatedCell, SpreadsheetFormulaCalculation, SpreadsheetFormulaValue, + MAX_SPREADSHEET_FORMULA_SPILL_CELLS, +}; + +use super::{ + calculation_storage_error, calculation_write_limit, insert_clear_write, CellWrite, + MAX_CALCULATION_WRITES, +}; + +pub(super) fn plan_writes( + document: &NativeOfficeDocument, + calculation: &SpreadsheetFormulaCalculation, +) -> UseResult>> { + let mut plans = BTreeMap::new(); + let mut planned_write_count = 0_usize; + for sheet in document + .root() + .children + .iter() + .filter(|node| node.node_type == OfficeNodeType::Worksheet) + { + let sheet_name = sheet.path.strip_prefix('/').ok_or_else(|| { + calculation_storage_error( + "use.office.spreadsheet_formula_storage_invalid", + format!("Worksheet path '{}' is invalid.", sheet.path), + ) + })?; + let part_name = sheet.format.get("part").cloned().ok_or_else(|| { + calculation_storage_error( + "use.office.spreadsheet_formula_storage_invalid", + format!("Worksheet '{}' has no source part.", sheet.path), + ) + })?; + let cells = worksheet_cells(sheet)?; + let mut writes = BTreeMap::new(); + plan_old_spill_cleanup(&cells, &mut writes)?; + for calculated in calculation + .cells + .iter() + .filter(|cell| cell.cell.sheet.eq_ignore_ascii_case(sheet_name)) + { + plan_calculated_cell(calculated, &cells, &mut writes)?; + } + validate_planned_writes(&cells, &writes)?; + if writes.len() > MAX_CALCULATION_WRITES { + return Err(calculation_storage_error( + "use.office.spreadsheet_formula_write_limit", + format!( + "Native formula recalculation writes at most {MAX_CALCULATION_WRITES} cells." + ), + ) + .with_detail("cells", writes.len())); + } + planned_write_count = planned_write_count + .checked_add(writes.len()) + .ok_or_else(calculation_write_limit)?; + if planned_write_count > MAX_CALCULATION_WRITES { + return Err(calculation_write_limit().with_detail("cells", planned_write_count)); + } + plans.insert(part_name, writes); + } + let planned_formulas = plans + .values() + .flat_map(BTreeMap::values) + .filter(|write| matches!(write, CellWrite::Formula { .. })) + .count(); + if planned_formulas != calculation.formula_count { + return Err(calculation_storage_error( + "use.office.spreadsheet_formula_storage_invalid", + "Calculation results do not match the workbook formula cells.", + ) + .with_detail("expectedFormulas", calculation.formula_count) + .with_detail("plannedFormulas", planned_formulas)); + } + Ok(plans) +} + +pub(super) fn worksheet_cells( + sheet: &DocumentNode, +) -> UseResult> { + let mut cells = BTreeMap::new(); + for cell in sheet + .children + .iter() + .filter(|node| node.node_type == OfficeNodeType::Row) + .flat_map(|row| &row.children) + .filter(|node| node.node_type == OfficeNodeType::Cell) + { + let reference = cell + .path + .rsplit_once('/') + .and_then(|(_, reference)| CellReference::parse(reference).ok()) + .ok_or_else(|| { + calculation_storage_error( + "use.office.spreadsheet_formula_storage_invalid", + format!("Spreadsheet cell path '{}' is invalid.", cell.path), + ) + })?; + if cells.insert(reference, cell).is_some() { + return Err(calculation_storage_error( + "use.office.spreadsheet_formula_storage_invalid", + format!("Worksheet contains duplicate cell '{}'.", reference.a1()), + )); + } + } + Ok(cells) +} + +fn plan_old_spill_cleanup( + cells: &BTreeMap, + writes: &mut BTreeMap, +) -> UseResult<()> { + for (anchor, cell) in cells { + let Some(reference) = cell.format.get("formulaRef") else { + continue; + }; + let range = CellRange::parse(reference).map_err(|error| { + calculation_storage_error( + "use.office.spreadsheet_formula_storage_invalid", + format!( + "Formula cell '{}' has invalid spill range '{reference}': {error}", + cell.path + ), + ) + })?; + if !range.contains(*anchor) { + return Err(calculation_storage_error( + "use.office.spreadsheet_formula_storage_invalid", + format!( + "Formula spill range '{}' does not contain anchor '{}'.", + range.a1(), + anchor.a1() + ), + )); + } + let spill_cells = range.cell_count()?; + if spill_cells > MAX_SPREADSHEET_FORMULA_SPILL_CELLS { + return Err(calculation_storage_error( + "use.office.spreadsheet_formula_spill_limit", + format!( + "Stored formula spill '{}' exceeds {MAX_SPREADSHEET_FORMULA_SPILL_CELLS} cells.", + range.a1() + ), + ) + .with_detail("cells", spill_cells)); + } + for row in range.start.row..=range.end.row { + for column in range.start.column..=range.end.column { + let reference = CellReference { column, row }; + if reference != *anchor { + insert_clear_write(writes, reference)?; + } + } + } + } + Ok(()) +} + +fn plan_calculated_cell( + calculated: &SpreadsheetFormulaCalculatedCell, + cells: &BTreeMap, + writes: &mut BTreeMap, +) -> UseResult<()> { + let anchor = CellReference { + column: calculated.cell.column, + row: calculated.cell.row, + }; + let cell = cells.get(&anchor).ok_or_else(|| { + calculation_storage_error( + "use.office.spreadsheet_formula_storage_invalid", + format!( + "Calculated formula cell '{}' is missing.", + calculated.cell.path() + ), + ) + })?; + let expression = cell.format.get("formula").cloned().ok_or_else(|| { + calculation_storage_error( + "use.office.spreadsheet_formula_storage_invalid", + format!("Calculated cell '{}' has no formula.", cell.path), + ) + })?; + if cell.format.get("formulaType").is_some_and(|value| { + !value.eq_ignore_ascii_case("normal") && !value.eq_ignore_ascii_case("array") + }) { + return Err(calculation_storage_error( + "use.office.spreadsheet_formula_storage_unsupported", + format!( + "Formula storage type '{}' at '{}' is not supported by native recalculation.", + cell.format.get("formulaType").map_or("", String::as_str), + cell.path + ), + )); + } + match &calculated.value { + SpreadsheetFormulaValue::Array { rows } => { + let spill = calculated.spill_range.as_deref().ok_or_else(|| { + calculation_storage_error( + "use.office.spreadsheet_formula_storage_invalid", + format!("Array result '{}' has no spill range.", cell.path), + ) + })?; + let range = CellRange::parse(spill)?; + let height = usize::try_from(range.end.row - range.start.row + 1) + .map_err(|_| calculation_write_limit())?; + let width = usize::try_from(range.end.column - range.start.column + 1) + .map_err(|_| calculation_write_limit())?; + if range.start != anchor + || rows.len() != height + || rows.iter().any(|row| row.len() != width) + { + return Err(calculation_storage_error( + "use.office.spreadsheet_formula_storage_invalid", + format!( + "Array result shape does not match spill range '{}' at '{}'.", + range.a1(), + cell.path + ), + )); + } + for (row_offset, row) in rows.iter().enumerate() { + for (column_offset, value) in row.iter().enumerate() { + require_scalar_value(value)?; + let reference = CellReference { + column: range.start.column + + u32::try_from(column_offset) + .map_err(|_| calculation_write_limit())?, + row: range.start.row + + u32::try_from(row_offset).map_err(|_| calculation_write_limit())?, + }; + let write = if reference == anchor { + CellWrite::Formula { + expression: expression.clone(), + value: value.clone(), + spill_range: Some(range.a1()), + } + } else { + CellWrite::Cached(value.clone()) + }; + insert_planned_write(writes, reference, write)?; + } + } + } + value => { + require_scalar_value(value)?; + if calculated.spill_range.is_some() { + return Err(calculation_storage_error( + "use.office.spreadsheet_formula_storage_invalid", + format!( + "Scalar result '{}' unexpectedly has a spill range.", + cell.path + ), + )); + } + insert_planned_write( + writes, + anchor, + CellWrite::Formula { + expression, + value: value.clone(), + spill_range: None, + }, + )?; + } + } + Ok(()) +} + +fn insert_planned_write( + writes: &mut BTreeMap, + reference: CellReference, + write: CellWrite, +) -> UseResult<()> { + match writes.get(&reference) { + None | Some(CellWrite::Clear) => { + if !writes.contains_key(&reference) { + let cells = writes + .len() + .checked_add(1) + .ok_or_else(calculation_write_limit)?; + if cells > MAX_CALCULATION_WRITES { + return Err(calculation_write_limit().with_detail("cells", cells)); + } + } + writes.insert(reference, write); + Ok(()) + } + Some(_) => Err(calculation_storage_error( + "use.office.spreadsheet_formula_spill_overlap", + format!( + "Calculated formula results overlap at '{}'.", + reference.a1() + ), + )), + } +} + +fn validate_planned_writes( + cells: &BTreeMap, + writes: &BTreeMap, +) -> UseResult<()> { + for (reference, write) in writes { + if matches!(write, CellWrite::Formula { .. }) { + continue; + } + if cells + .get(reference) + .is_some_and(|cell| cell.format.contains_key("formula")) + { + return Err(calculation_storage_error( + "use.office.spreadsheet_formula_spill_overlap", + format!( + "Calculated spill at '{}' overlaps another formula cell.", + reference.a1() + ), + )); + } + } + Ok(()) +} + +fn require_scalar_value(value: &SpreadsheetFormulaValue) -> UseResult<()> { + if matches!(value, SpreadsheetFormulaValue::Array { .. }) { + return Err(calculation_storage_error( + "use.office.spreadsheet_formula_storage_invalid", + "Nested Spreadsheet formula arrays cannot be written to OOXML cells.", + )); + } + Ok(()) +} diff --git a/crates/office/src/editor/spreadsheet/formula/write.rs b/crates/office/src/editor/spreadsheet/formula/write.rs new file mode 100644 index 00000000..254fc0a0 --- /dev/null +++ b/crates/office/src/editor/spreadsheet/formula/write.rs @@ -0,0 +1,354 @@ +use std::collections::BTreeMap; + +use a3s_use_core::UseResult; + +use crate::spreadsheet_reference::CellReference; +use crate::xml_edit::{insert_child, IndexedXmlElement, XmlPatch}; +use crate::{NativeOfficePackage, SpreadsheetFormulaValue}; + +use super::super::{ + escape_attribute, expanded_element, indexed_cells_in_row, indexed_rows, prefix, qualified, + remove_calculation_chain, +}; +use super::{calculation_storage_error, CellWrite}; + +pub(super) fn apply_cell_writes( + part: &crate::LosslessXmlPart, + sheet_data: &IndexedXmlElement, + writes: &BTreeMap, +) -> UseResult> { + let mut by_row = BTreeMap::>::new(); + for (reference, write) in writes { + by_row + .entry(reference.row) + .or_default() + .push((*reference, write)); + } + if sheet_data.empty { + let rows = by_row + .into_iter() + .filter_map(|(row_number, writes)| { + let cells = writes + .into_iter() + .filter_map(|(reference, write)| { + new_cell_fragment(prefix(&sheet_data.qualified_name), reference, write) + .transpose() + }) + .collect::>(); + match cells { + Ok(cells) if cells.is_empty() => None, + Ok(cells) => { + let tag = qualified(prefix(&sheet_data.qualified_name), "row"); + Some(Ok(format!("<{tag} r=\"{row_number}\">{cells}"))) + } + Err(error) => Some(Err(error)), + } + }) + .collect::>()?; + return insert_child(part, sheet_data, rows); + } + + let rows = indexed_rows(sheet_data); + let row_map = rows.iter().copied().collect::>(); + let mut patches = Vec::new(); + let mut insertions = BTreeMap::>::new(); + for (row_number, row_writes) in by_row { + let Some(row) = row_map.get(&row_number).copied() else { + let cells = row_writes + .into_iter() + .filter_map(|(reference, write)| { + new_cell_fragment(prefix(&sheet_data.qualified_name), reference, write) + .transpose() + }) + .collect::>()?; + if cells.is_empty() { + continue; + } + let tag = qualified(prefix(&sheet_data.qualified_name), "row"); + let fragment = format!("<{tag} r=\"{row_number}\">{cells}"); + let position = rows + .iter() + .find(|(existing, _)| *existing > row_number) + .map_or(sheet_data.content_range.end, |(_, next)| { + next.full_range.start + }); + insertions + .entry(position) + .or_default() + .push((row_number, 0, fragment)); + continue; + }; + if row.empty { + let cells = row_writes + .into_iter() + .filter_map(|(reference, write)| { + new_cell_fragment(prefix(&row.qualified_name), reference, write).transpose() + }) + .collect::>()?; + if !cells.is_empty() { + patches.push(XmlPatch::new( + row.full_range.clone(), + expanded_element(row, &cells), + )); + } + continue; + } + let cells = indexed_cells_in_row(row_number, row); + for (reference, write) in row_writes { + if let Some((_, cell)) = cells + .iter() + .find(|(existing, _)| existing.column == reference.column) + { + let replacement = existing_cell_fragment(part, cell, reference, write)?; + patches.push(XmlPatch::new( + cell.full_range.clone(), + replacement.unwrap_or_default(), + )); + continue; + } + let Some(fragment) = new_cell_fragment(prefix(&row.qualified_name), reference, write)? + else { + continue; + }; + let position = cells + .iter() + .find(|(existing, _)| existing.column > reference.column) + .map(|(_, next)| next.full_range.start) + .or_else(|| { + row.children + .iter() + .find(|child| child.local_name != "c") + .map(|child| child.full_range.start) + }) + .unwrap_or(row.content_range.end); + insertions + .entry(position) + .or_default() + .push((row_number, reference.column, fragment)); + } + } + for (position, mut fragments) in insertions { + fragments.sort_by_key(|(row, column, _)| (*row, *column)); + patches.push(XmlPatch::new( + position..position, + fragments + .into_iter() + .map(|(_, _, fragment)| fragment) + .collect::(), + )); + } + crate::xml_edit::apply_patches(part, patches) +} + +fn new_cell_fragment( + namespace_prefix: Option<&str>, + reference: CellReference, + write: &CellWrite, +) -> UseResult> { + if matches!(write, CellWrite::Clear) { + return Ok(None); + } + let tag = qualified(namespace_prefix, "c"); + let (value_type, content) = owned_cell_content(namespace_prefix, None, write)?; + let value_type = value_type.map_or_else(String::new, |value_type| { + format!(" t=\"{}\"", escape_attribute(value_type)) + }); + Ok(Some(format!( + "<{tag} r=\"{}\"{value_type}>{content}", + reference.a1() + ))) +} + +fn existing_cell_fragment( + part: &crate::LosslessXmlPart, + cell: &IndexedXmlElement, + reference: CellReference, + write: &CellWrite, +) -> UseResult>> { + let preserved = preserved_cell_content(part, cell)?; + let mut attributes = cell.qualified_attributes.clone(); + attributes.insert("r".into(), reference.a1()); + let (value_type, owned) = + owned_cell_content(prefix(&cell.qualified_name), cell.child("f", 1), write)?; + if let Some(value_type) = value_type { + attributes.insert("t".into(), value_type.to_string()); + } else { + attributes.remove("t"); + } + if matches!(write, CellWrite::Clear) + && attributes.len() == 1 + && attributes.contains_key("r") + && preserved.iter().all(u8::is_ascii_whitespace) + { + return Ok(None); + } + let attributes = attributes + .into_iter() + .map(|(name, value)| format!(" {name}=\"{}\"", escape_attribute(&value))) + .collect::(); + let mut content = owned.into_bytes(); + content.extend_from_slice(&preserved); + if content.is_empty() { + return Ok(Some( + format!("<{}{attributes}/>", cell.qualified_name).into_bytes(), + )); + } + let mut output = format!("<{}{attributes}>", cell.qualified_name).into_bytes(); + output.extend_from_slice(&content); + output.extend_from_slice(format!("", cell.qualified_name).as_bytes()); + Ok(Some(output)) +} + +fn preserved_cell_content( + part: &crate::LosslessXmlPart, + cell: &IndexedXmlElement, +) -> UseResult> { + if cell.empty { + return Ok(Vec::new()); + } + let bytes = part.parse_bytes(); + let mut output = Vec::new(); + let mut cursor = cell.content_range.start; + for child in &cell.children { + if matches!(child.local_name.as_str(), "f" | "v" | "is") { + output.extend_from_slice(bytes.get(cursor..child.full_range.start).ok_or_else( + || { + calculation_storage_error( + "use.office.spreadsheet_formula_storage_invalid", + "Spreadsheet cell child range is invalid.", + ) + }, + )?); + cursor = child.full_range.end; + } + } + output.extend_from_slice(bytes.get(cursor..cell.content_range.end).ok_or_else(|| { + calculation_storage_error( + "use.office.spreadsheet_formula_storage_invalid", + "Spreadsheet cell content range is invalid.", + ) + })?); + Ok(output) +} + +fn owned_cell_content( + namespace_prefix: Option<&str>, + existing_formula: Option<&IndexedXmlElement>, + write: &CellWrite, +) -> UseResult<(Option<&'static str>, String)> { + match write { + CellWrite::Clear => Ok((None, String::new())), + CellWrite::Cached(value) => cached_value_content(namespace_prefix, value), + CellWrite::Formula { + expression, + value, + spill_range, + } => { + let formula_tag = qualified(namespace_prefix, "f"); + let mut attributes = existing_formula + .map(|formula| formula.qualified_attributes.clone()) + .unwrap_or_default(); + let existing_type = attributes.get("t").cloned(); + attributes.remove("t"); + attributes.remove("ref"); + if let Some(spill_range) = spill_range { + attributes.insert("t".into(), "array".into()); + attributes.insert("ref".into(), spill_range.clone()); + } else if existing_type.is_some_and(|value| value.eq_ignore_ascii_case("normal")) { + attributes.insert("t".into(), "normal".into()); + } + let attributes = attributes + .into_iter() + .map(|(name, value)| format!(" {name}=\"{}\"", escape_attribute(&value))) + .collect::(); + let formula = format!( + "<{formula_tag}{attributes}>{}", + crate::xml_edit::escape_text(expression) + ); + let (value_type, value) = cached_value_content(namespace_prefix, value)?; + Ok((value_type, format!("{formula}{value}"))) + } + } +} + +fn cached_value_content( + namespace_prefix: Option<&str>, + value: &SpreadsheetFormulaValue, +) -> UseResult<(Option<&'static str>, String)> { + let value_tag = qualified(namespace_prefix, "v"); + let (value_type, value) = match value { + SpreadsheetFormulaValue::Blank => (None, String::new()), + SpreadsheetFormulaValue::Number { value } => (None, value.clone()), + SpreadsheetFormulaValue::Text { value } => (Some("str"), value.clone()), + SpreadsheetFormulaValue::Boolean { value } => { + (Some("b"), if *value { "1".into() } else { "0".into() }) + } + SpreadsheetFormulaValue::Error { error } => (Some("e"), error.as_str().into()), + SpreadsheetFormulaValue::Array { .. } => { + return Err(calculation_storage_error( + "use.office.spreadsheet_formula_storage_invalid", + "Nested Spreadsheet formula arrays cannot be written to OOXML cells.", + )) + } + }; + Ok(( + value_type, + format!( + "<{value_tag}>{}", + crate::xml_edit::escape_text(&value) + ), + )) +} + +pub(super) fn mark_workbook_calculated(package: &mut NativeOfficePackage) -> UseResult<()> { + remove_calculation_chain(package)?; + let workbook = package.xml_part("xl/workbook.xml")?; + let index = crate::xml_edit::index_xml(&workbook)?; + let edited = if let Some(calc) = index.child("calcPr", 1) { + let mut attributes = calc.qualified_attributes.clone(); + attributes.insert("calcMode".into(), "auto".into()); + attributes.insert("calcCompleted".into(), "1".into()); + attributes.insert("fullCalcOnLoad".into(), "0".into()); + attributes.insert("forceFullCalc".into(), "0".into()); + let attributes = attributes + .into_iter() + .map(|(name, value)| format!(" {name}=\"{}\"", escape_attribute(&value))) + .collect::(); + let terminator = if calc.empty { "/>" } else { ">" }; + crate::xml_edit::apply_patches( + &workbook, + vec![XmlPatch::new( + calc.start_tag_range.clone(), + format!("<{}{attributes}{terminator}", calc.qualified_name), + )], + )? + } else { + let tag = qualified(prefix(&index.qualified_name), "calcPr"); + let fragment = format!( + "<{tag} calcId=\"0\" calcMode=\"auto\" calcCompleted=\"1\" fullCalcOnLoad=\"0\" forceFullCalc=\"0\"/>" + ); + let insertion = index + .children + .iter() + .find(|child| { + matches!( + child.local_name.as_str(), + "oleSize" + | "customWorkbookViews" + | "pivotCaches" + | "smartTagPr" + | "smartTagTypes" + | "webPublishing" + | "fileRecoveryPr" + | "webPublishObjects" + | "extLst" + ) + }) + .map_or(index.content_range.end, |child| child.full_range.start); + crate::xml_edit::apply_patches( + &workbook, + vec![XmlPatch::new(insertion..insertion, fragment)], + )? + }; + package.set_part("xl/workbook.xml", edited) +} diff --git a/crates/office/src/editor/spreadsheet/import.rs b/crates/office/src/editor/spreadsheet/import.rs new file mode 100644 index 00000000..d0975a6b --- /dev/null +++ b/crates/office/src/editor/spreadsheet/import.rs @@ -0,0 +1,514 @@ +use std::collections::{BTreeMap, BTreeSet}; + +use a3s_use_core::UseResult; + +use super::{ + editor_error, expanded_element, indexed_cells, indexed_cells_in_row, indexed_rows, + mark_workbook_for_recalculation, prefix, qualified, update_dimension, +}; +use crate::editor::{ + NativeSpreadsheetAutoFilter, NativeSpreadsheetCellFormat, NativeSpreadsheetDelimitedImport, + NativeSpreadsheetFrozenPane, NativeSpreadsheetImportResult, + MAX_NATIVE_SPREADSHEET_IMPORT_CELLS, +}; +use crate::semantic::{NativeOfficeDocument, OfficeNodeType}; +use crate::spreadsheet_reference::{CellRange, CellReference, MAX_COLUMNS, MAX_ROWS}; +use crate::xml_edit::{apply_patches, index_xml, insert_child, IndexedXmlElement, XmlPatch}; +use crate::{DocumentKind, LosslessXmlPart, NativeOfficePackage}; + +mod parse; + +pub(super) fn apply( + package: &mut NativeOfficePackage, + requested_sheet: &str, + request: &NativeSpreadsheetDelimitedImport, +) -> UseResult { + let sheet = resolve_sheet(package, requested_sheet)?; + let start = request.validate()?; + let parsed = parse::parse(request, workbook_uses_1904_date_system(package)?)?; + if parsed.rows.is_empty() { + return Ok(NativeSpreadsheetImportResult { + path: sheet.path.clone(), + sheet: sheet.path, + start_cell: start.a1(), + range: None, + format: request.format, + row_count: 0, + column_count: 0, + header: request.header, + changed: false, + filter_path: None, + freeze_path: None, + }); + } + let range = target_range(start, parsed.rows.len(), parsed.max_columns)?; + let worksheet = package.xml_part(&sheet.part)?; + let root = index_xml(&worksheet)?; + let sheet_data = root + .descendant("sheetData") + .ok_or_else(|| import_part_error(&sheet.part, "has no sheetData element"))?; + let edited = write_rows(package, &worksheet, sheet_data, start, &parsed.rows)?; + let edited = update_dimension(&sheet.part, edited)?; + package.set_part(&sheet.part, edited)?; + mark_workbook_for_recalculation(package)?; + + let (filter_path, freeze_path) = if request.header { + let filter = NativeSpreadsheetAutoFilter::new(range.a1()); + let snapshot = NativeOfficeDocument::from_package(package.clone())?; + let existing = snapshot + .get(&sheet.path, 1)? + .children + .into_iter() + .find(|node| node.node_type == OfficeNodeType::AutoFilter); + let filter_path = if let Some(existing) = existing { + super::auto_filter::set(package, &existing.path, &filter)? + } else { + super::auto_filter::add(package, &sheet.path, &filter)? + }; + let top_row = start + .row + .checked_add(1) + .filter(|row| *row <= MAX_ROWS) + .ok_or_else(|| { + import_error( + "use.office.spreadsheet_import_row_limit", + "A header import cannot freeze below Excel's final worksheet row.", + ) + })?; + let top_left = CellReference { + column: start.column, + row: top_row, + } + .a1(); + let freeze_path = super::view::set( + package, + &sheet.path, + &NativeSpreadsheetFrozenPane::new(start.row, 0, top_left), + )?; + (Some(filter_path), Some(freeze_path)) + } else { + (None, None) + }; + + let range_name = range.a1(); + Ok(NativeSpreadsheetImportResult { + path: format!("{}/{}", sheet.path, range_name), + sheet: sheet.path, + start_cell: start.a1(), + range: Some(range_name), + format: request.format, + row_count: parsed.rows.len(), + column_count: parsed.max_columns, + header: request.header, + changed: true, + filter_path, + freeze_path, + }) +} + +struct ResolvedSheet { + path: String, + part: String, +} + +fn resolve_sheet(package: &NativeOfficePackage, requested: &str) -> UseResult { + if package.kind() != DocumentKind::Spreadsheet { + return Err(import_error( + "use.office.mutation_type_unsupported", + "Delimited import is available only for Spreadsheet documents.", + )); + } + if !requested.starts_with('/') + || requested.len() < 2 + || requested.trim_start_matches('/').contains('/') + || requested.chars().any(char::is_control) + { + return Err(import_error( + "use.office.mutation_path_unsupported", + "Spreadsheet import requires a worksheet path such as /Sheet1.", + )); + } + let snapshot = NativeOfficeDocument::from_package(package.clone())?; + let sheet = snapshot + .root() + .children + .iter() + .find(|node| { + node.node_type == OfficeNodeType::Worksheet && node.path.eq_ignore_ascii_case(requested) + }) + .ok_or_else(|| { + import_error( + "use.office.node_not_found", + format!("Office semantic path '{requested}' does not exist."), + ) + })?; + Ok(ResolvedSheet { + path: sheet.path.clone(), + part: sheet.format.get("part").cloned().ok_or_else(|| { + import_error( + "use.office.spreadsheet_sheet_invalid", + format!("Worksheet '{}' has no source part.", sheet.path), + ) + })?, + }) +} + +fn target_range(start: CellReference, rows: usize, columns: usize) -> UseResult { + let row_count = u32::try_from(rows).map_err(|_| row_limit())?; + let column_count = u32::try_from(columns).map_err(|_| column_limit())?; + let end_row = start + .row + .checked_add(row_count.saturating_sub(1)) + .filter(|row| *row <= MAX_ROWS) + .ok_or_else(row_limit)?; + let end_column = start + .column + .checked_add(column_count.saturating_sub(1)) + .filter(|column| *column <= MAX_COLUMNS) + .ok_or_else(column_limit)?; + let cells = rows.checked_mul(columns).ok_or_else(cell_limit)?; + if cells > MAX_NATIVE_SPREADSHEET_IMPORT_CELLS { + return Err(cell_limit().with_detail("cells", cells)); + } + Ok(CellRange { + start, + end: CellReference { + column: end_column, + row: end_row, + }, + }) +} + +fn write_rows( + package: &mut NativeOfficePackage, + worksheet: &LosslessXmlPart, + sheet_data: &IndexedXmlElement, + start: CellReference, + source_rows: &[Vec], +) -> UseResult> { + let rows = indexed_rows(sheet_data); + let row_map = rows.iter().copied().collect::>(); + let existing_cells = indexed_cells(sheet_data) + .into_iter() + .map(|(reference, _, cell)| (reference, cell)) + .collect::>(); + let mut date_base_styles = BTreeSet::new(); + for (row_offset, fields) in source_rows.iter().enumerate() { + for (column_offset, field) in fields.iter().enumerate() { + if !field.value().is_some_and(|(_, date)| date) { + continue; + } + let reference = source_reference(start, row_offset, column_offset)?; + let base = existing_cells + .get(&reference) + .map(|cell| super::style::cell_style_index(cell)) + .transpose()? + .unwrap_or(0); + date_base_styles.insert(base); + } + } + let date_styles = if date_base_styles.is_empty() { + BTreeMap::new() + } else { + super::style::derived_cell_style_indexes( + package, + &date_base_styles, + &NativeSpreadsheetCellFormat { + number_format: Some("date".into()), + ..NativeSpreadsheetCellFormat::default() + }, + )? + }; + + if sheet_data.empty { + let prefix = prefix(&sheet_data.qualified_name); + let row_tag = qualified(prefix, "row"); + let mut rows = String::new(); + for (row_offset, fields) in source_rows.iter().enumerate() { + let cells = new_cells(prefix, start, row_offset, fields, &date_styles)?; + if cells.is_empty() { + continue; + } + let row_number = start + .row + .checked_add(u32::try_from(row_offset).map_err(|_| row_limit())?) + .ok_or_else(row_limit)?; + rows.push_str(&format!( + "<{row_tag} r=\"{row_number}\">{cells}" + )); + } + return insert_child(worksheet, sheet_data, rows); + } + + let mut patches = Vec::new(); + let mut insertions = BTreeMap::>::new(); + let sheet_prefix = prefix(&sheet_data.qualified_name); + for (row_offset, fields) in source_rows.iter().enumerate() { + let row_number = start + .row + .checked_add(u32::try_from(row_offset).map_err(|_| row_limit())?) + .ok_or_else(row_limit)?; + let Some(row) = row_map.get(&row_number).copied() else { + let cells = new_cells(sheet_prefix, start, row_offset, fields, &date_styles)?; + if cells.is_empty() { + continue; + } + let row_tag = qualified(sheet_prefix, "row"); + let fragment = format!("<{row_tag} r=\"{row_number}\">{cells}"); + let position = rows + .iter() + .find(|(existing, _)| *existing > row_number) + .map_or(sheet_data.content_range.end, |(_, next)| { + next.full_range.start + }); + insertions + .entry(position) + .or_default() + .push((row_number, 0, fragment)); + continue; + }; + if row.empty { + let cells = new_cells( + prefix(&row.qualified_name), + start, + row_offset, + fields, + &date_styles, + )?; + if !cells.is_empty() { + patches.push(XmlPatch::new( + row.full_range.clone(), + expanded_element(row, &cells), + )); + } + continue; + } + let cells = indexed_cells_in_row(row_number, row); + for (column_offset, field) in fields.iter().enumerate() { + let reference = source_reference(start, row_offset, column_offset)?; + if let Some((_, cell)) = cells + .iter() + .find(|(existing, _)| existing.column == reference.column) + { + patches.push(XmlPatch::new( + cell.full_range.clone(), + existing_cell_fragment(worksheet, cell, reference, field, &date_styles)?, + )); + } else if !field.is_empty() { + let position = cells + .iter() + .find(|(existing, _)| existing.column > reference.column) + .map(|(_, next)| next.full_range.start) + .or_else(|| { + row.children + .iter() + .find(|child| child.local_name != "c") + .map(|child| child.full_range.start) + }) + .unwrap_or(row.content_range.end); + insertions.entry(position).or_default().push(( + row_number, + reference.column, + new_cell(prefix(&row.qualified_name), reference, field, &date_styles)?, + )); + } + } + } + for (position, mut fragments) in insertions { + fragments.sort_by_key(|(row, column, _)| (*row, *column)); + patches.push(XmlPatch::new( + position..position, + fragments + .into_iter() + .map(|(_, _, fragment)| fragment) + .collect::(), + )); + } + apply_patches(worksheet, patches) +} + +fn new_cells( + prefix: Option<&str>, + start: CellReference, + row_offset: usize, + fields: &[parse::ParsedField], + date_styles: &BTreeMap, +) -> UseResult { + fields + .iter() + .enumerate() + .filter(|(_, field)| !field.is_empty()) + .map(|(column_offset, field)| { + let reference = source_reference(start, row_offset, column_offset)?; + new_cell(prefix, reference, field, date_styles) + }) + .collect() +} + +fn new_cell( + prefix: Option<&str>, + reference: CellReference, + field: &parse::ParsedField, + date_styles: &BTreeMap, +) -> UseResult { + let (value, date) = field.value().ok_or_else(import_cell_invalid)?; + let tag = qualified(prefix, "c"); + let (value_type, content) = super::cell_content(prefix, value); + let style = if date { + format!(" s=\"{}\"", required_date_style(date_styles, 0)?) + } else { + String::new() + }; + let value_type = value_type + .map(|value_type| format!(" t=\"{value_type}\"")) + .unwrap_or_default(); + Ok(format!( + "<{tag} r=\"{}\"{style}{value_type}>{content}", + reference.a1() + )) +} + +fn existing_cell_fragment( + part: &LosslessXmlPart, + cell: &IndexedXmlElement, + reference: CellReference, + field: &parse::ParsedField, + date_styles: &BTreeMap, +) -> UseResult { + let mut attributes = cell.qualified_attributes.clone(); + attributes.insert("r".into(), reference.a1()); + let mut content = String::new(); + if let Some((value, date)) = field.value() { + let (value_type, value_content) = super::cell_content(prefix(&cell.qualified_name), value); + if let Some(value_type) = value_type { + attributes.insert("t".into(), value_type.into()); + } else { + attributes.remove("t"); + } + if date { + let base = super::style::cell_style_index(cell)?; + attributes.insert( + "s".into(), + required_date_style(date_styles, base)?.to_string(), + ); + } + content.push_str(&value_content); + } else { + attributes.remove("t"); + } + for child in &cell.children { + let owned_value = child.namespace == cell.namespace + && matches!(child.local_name.as_str(), "f" | "v" | "is"); + if !owned_value { + let bytes = &part.parse_bytes()[child.full_range.clone()]; + content.push_str(std::str::from_utf8(bytes).map_err(|error| { + import_error( + "use.office.spreadsheet_import_cell_invalid", + format!("Spreadsheet cell extension content is not UTF-8: {error}"), + ) + })?); + } + } + let attributes = attributes + .into_iter() + .map(|(name, value)| format!(" {name}=\"{}\"", crate::xml_edit::escape_attribute(&value))) + .collect::(); + let tag = qualified(prefix(&cell.qualified_name), "c"); + Ok(format!("<{tag}{attributes}>{content}")) +} + +fn source_reference( + start: CellReference, + row_offset: usize, + column_offset: usize, +) -> UseResult { + Ok(CellReference { + column: start + .column + .checked_add(u32::try_from(column_offset).map_err(|_| column_limit())?) + .ok_or_else(column_limit)?, + row: start + .row + .checked_add(u32::try_from(row_offset).map_err(|_| row_limit())?) + .ok_or_else(row_limit)?, + }) +} + +fn required_date_style(styles: &BTreeMap, base: usize) -> UseResult { + styles.get(&base).copied().ok_or_else(|| { + import_error( + "use.office.spreadsheet_styles_invalid", + format!("Spreadsheet import could not derive a date style from style {base}."), + ) + }) +} + +fn workbook_uses_1904_date_system(package: &NativeOfficePackage) -> UseResult { + let workbook = package.xml_part("xl/workbook.xml")?; + let root = index_xml(&workbook)?; + let properties = root + .children + .iter() + .filter(|child| child.local_name == "workbookPr" && child.namespace == root.namespace) + .collect::>(); + if properties.len() > 1 { + return Err(import_error( + "use.office.spreadsheet_import_date_system_invalid", + "Spreadsheet workbook contains multiple workbookPr elements.", + )); + } + match properties + .first() + .and_then(|properties| properties.attributes.get("date1904")) + .map(String::as_str) + { + None | Some("0" | "false") => Ok(false), + Some("1" | "true") => Ok(true), + Some(value) => Err(import_error( + "use.office.spreadsheet_import_date_system_invalid", + format!("Spreadsheet workbook has invalid date1904='{value}'."), + )), + } +} + +fn row_limit() -> a3s_use_core::UseError { + import_error( + "use.office.spreadsheet_import_row_limit", + format!("Spreadsheet import cannot exceed Excel's {MAX_ROWS} rows."), + ) +} + +fn column_limit() -> a3s_use_core::UseError { + import_error( + "use.office.spreadsheet_import_column_limit", + format!("Spreadsheet import cannot exceed Excel's {MAX_COLUMNS} columns."), + ) +} + +fn cell_limit() -> a3s_use_core::UseError { + import_error( + "use.office.spreadsheet_import_cell_count_limit", + format!( + "Spreadsheet import accepts at most {MAX_NATIVE_SPREADSHEET_IMPORT_CELLS} rectangular target cells." + ), + ) +} + +fn import_cell_invalid() -> a3s_use_core::UseError { + import_error( + "use.office.spreadsheet_import_cell_invalid", + "Spreadsheet import attempted to materialize an empty source field.", + ) +} + +fn import_part_error(part: &str, reason: &str) -> a3s_use_core::UseError { + import_error( + "use.office.spreadsheet_import_part_invalid", + format!("Spreadsheet worksheet part '{part}' {reason}."), + ) + .with_detail("part", part) +} + +fn import_error(code: &str, message: impl Into) -> a3s_use_core::UseError { + editor_error(code, message) +} diff --git a/crates/office/src/editor/spreadsheet/import/parse.rs b/crates/office/src/editor/spreadsheet/import/parse.rs new file mode 100644 index 00000000..8704abf0 --- /dev/null +++ b/crates/office/src/editor/spreadsheet/import/parse.rs @@ -0,0 +1,435 @@ +use a3s_use_core::UseResult; + +use crate::editor::{ + NativeSpreadsheetDelimitedImport, SpreadsheetCellValue, MAX_NATIVE_SPREADSHEET_IMPORT_CELLS, +}; + +const MAX_CELL_UTF16_UNITS: usize = 32_767; + +#[derive(Debug)] +pub(super) struct ParsedImport { + pub(super) rows: Vec>, + pub(super) max_columns: usize, +} + +#[derive(Debug)] +pub(super) enum ParsedField { + Empty, + Value { + value: SpreadsheetCellValue, + date: bool, + }, +} + +impl ParsedField { + pub(super) fn is_empty(&self) -> bool { + matches!(self, Self::Empty) + } + + pub(super) fn value(&self) -> Option<(&SpreadsheetCellValue, bool)> { + match self { + Self::Empty => None, + Self::Value { value, date } => Some((value, *date)), + } + } +} + +pub(super) fn parse( + request: &NativeSpreadsheetDelimitedImport, + date_1904: bool, +) -> UseResult { + let rows = parse_delimited(&request.content, request.format.delimiter())?; + let max_columns = rows.iter().map(Vec::len).max().unwrap_or(0); + let rows = rows + .into_iter() + .map(|row| { + row.into_iter() + .map(|field| parse_field(field, date_1904)) + .collect::>>() + }) + .collect::>>()?; + Ok(ParsedImport { rows, max_columns }) +} + +fn parse_delimited(content: &str, delimiter: char) -> UseResult>> { + let content = content.strip_prefix('\u{feff}').unwrap_or(content); + if content.is_empty() { + return Ok(Vec::new()); + } + let mut rows = Vec::new(); + let mut row = Vec::new(); + let mut field = String::new(); + let mut in_quotes = false; + let mut field_started = false; + let mut quote_closed = false; + let mut fields = 0_usize; + let mut maximum_columns = 0_usize; + let mut characters = content.chars().peekable(); + while let Some(character) = characters.next() { + if in_quotes { + if character == '"' { + if characters.peek() == Some(&'"') { + field.push('"'); + characters.next(); + } else { + in_quotes = false; + quote_closed = true; + } + } else { + field.push(character); + } + continue; + } + if character == '"' && !field_started { + in_quotes = true; + field_started = true; + } else if character == delimiter { + row.push(std::mem::take(&mut field)); + fields = checked_field_count(fields)?; + field_started = false; + quote_closed = false; + } else if character == '\r' { + row.push(std::mem::take(&mut field)); + fields = checked_field_count(fields)?; + push_row(&mut rows, std::mem::take(&mut row), &mut maximum_columns)?; + field_started = false; + quote_closed = false; + if characters.peek() == Some(&'\n') { + characters.next(); + } + } else if character == '\n' { + row.push(std::mem::take(&mut field)); + fields = checked_field_count(fields)?; + push_row(&mut rows, std::mem::take(&mut row), &mut maximum_columns)?; + field_started = false; + quote_closed = false; + } else if character == '"' { + return Err(delimited_invalid( + &rows, + &row, + "A quote may appear only at the beginning of a delimited field.", + )); + } else if quote_closed { + return Err(delimited_invalid( + &rows, + &row, + "Only a delimiter or record boundary may follow a closing quote.", + )); + } else { + field.push(character); + field_started = true; + } + } + if in_quotes { + return Err(delimited_invalid( + &rows, + &row, + "Delimited input ended before a quoted field was closed.", + )); + } + if field_started || !row.is_empty() { + row.push(field); + checked_field_count(fields)?; + push_row(&mut rows, row, &mut maximum_columns)?; + } + Ok(rows) +} + +fn checked_field_count(fields: usize) -> UseResult { + let fields = fields.checked_add(1).ok_or_else(import_cell_count_limit)?; + if fields > MAX_NATIVE_SPREADSHEET_IMPORT_CELLS { + return Err(import_cell_count_limit().with_detail("fields", fields)); + } + Ok(fields) +} + +fn push_row( + rows: &mut Vec>, + row: Vec, + maximum_columns: &mut usize, +) -> UseResult<()> { + *maximum_columns = (*maximum_columns).max(row.len()); + let row_count = rows + .len() + .checked_add(1) + .ok_or_else(import_cell_count_limit)?; + let cells = row_count + .checked_mul(*maximum_columns) + .ok_or_else(import_cell_count_limit)?; + if cells > MAX_NATIVE_SPREADSHEET_IMPORT_CELLS { + return Err(import_cell_count_limit().with_detail("cells", cells)); + } + rows.push(row); + Ok(()) +} + +fn parse_field(value: String, date_1904: bool) -> UseResult { + if value.is_empty() { + return Ok(ParsedField::Empty); + } + validate_cell_text(&value)?; + let (value, date) = if let Some(expression) = value.strip_prefix('=') { + ( + super::super::normalize_cell_value(&SpreadsheetCellValue::Formula { + expression: expression.to_string(), + })?, + false, + ) + } else if let Some(number) = normalize_number(&value) { + ( + super::super::normalize_cell_value(&SpreadsheetCellValue::Number { value: number })?, + false, + ) + } else if let Some(serial) = parse_iso_date(&value, date_1904) { + ( + super::super::normalize_cell_value(&SpreadsheetCellValue::Number { value: serial })?, + true, + ) + } else if value.eq_ignore_ascii_case("true") { + (SpreadsheetCellValue::Boolean { value: true }, false) + } else if value.eq_ignore_ascii_case("false") { + (SpreadsheetCellValue::Boolean { value: false }, false) + } else { + (SpreadsheetCellValue::Text { value }, false) + }; + Ok(ParsedField::Value { value, date }) +} + +fn validate_cell_text(value: &str) -> UseResult<()> { + let units = value.encode_utf16().count(); + if units > MAX_CELL_UTF16_UNITS { + return Err(import_error( + "use.office.spreadsheet_import_cell_limit", + format!( + "Spreadsheet import fields cannot exceed {MAX_CELL_UTF16_UNITS} UTF-16 code units." + ), + ) + .with_detail("utf16Units", units)); + } + if let Some(character) = value.chars().find(|character| { + !matches!( + u32::from(*character), + 0x9 | 0xA | 0xD | 0x20..=0xD7FF | 0xE000..=0xFFFD | 0x10000..=0x10FFFF + ) + }) { + return Err(import_error( + "use.office.spreadsheet_import_cell_invalid", + format!( + "Spreadsheet import field contains XML-forbidden character U+{:04X}.", + u32::from(character) + ), + )); + } + Ok(()) +} + +fn normalize_number(value: &str) -> Option { + let parsed = value.parse::().ok().filter(|value| value.is_finite()); + if let Some(parsed) = parsed { + return Some(if is_canonical_number(value) { + value.to_string() + } else { + parsed.to_string() + }); + } + + let mut candidate = value.trim(); + let parenthesized = candidate.starts_with('(') && candidate.ends_with(')'); + if parenthesized { + candidate = &candidate[1..candidate.len() - 1]; + } + candidate = candidate.trim(); + if let Some(rest) = candidate.strip_prefix('$') { + candidate = rest; + } + let normalized = candidate.replace(',', ""); + let parsed = normalized + .parse::() + .ok() + .filter(|value| value.is_finite())?; + let parsed = if parenthesized { -parsed } else { parsed }; + Some(parsed.to_string()) +} + +fn is_canonical_number(value: &str) -> bool { + let value = value.strip_prefix('-').unwrap_or(value); + if value.is_empty() { + return false; + } + let mut exponent_parts = value.split(['e', 'E']); + let mantissa = exponent_parts.next().unwrap_or_default(); + let exponent = exponent_parts.next(); + if exponent_parts.next().is_some() + || exponent.is_some_and(|value| { + let digits = value + .strip_prefix('+') + .or_else(|| value.strip_prefix('-')) + .unwrap_or(value); + digits.is_empty() || !digits.bytes().all(|byte| byte.is_ascii_digit()) + }) + { + return false; + } + let mut decimal_parts = mantissa.split('.'); + let whole = decimal_parts.next().unwrap_or_default(); + let fraction = decimal_parts.next(); + if decimal_parts.next().is_some() { + return false; + } + let whole_valid = whole.bytes().all(|byte| byte.is_ascii_digit()); + let fraction_valid = fraction.is_none_or(|digits| { + digits.bytes().all(|byte| byte.is_ascii_digit()) + && (!whole.is_empty() || !digits.is_empty()) + }); + whole_valid && fraction_valid && (!whole.is_empty() || fraction.is_some()) +} + +fn parse_iso_date(value: &str, date_1904: bool) -> Option { + let (date, time) = if value.len() == 10 { + (value, None) + } else if value.len() >= 19 && matches!(value.as_bytes().get(10), Some(b'T' | b' ')) { + (&value[..10], Some(&value[11..])) + } else { + return None; + }; + let (year, month, day) = parse_date(date)?; + let (hour, minute, second, millis) = match time { + None => (0, 0, 0, 0), + Some(time) => parse_time(time)?, + }; + if year < 100 { + return None; + } + let baseline = days_from_civil(1899, 12, 30); + let mut serial = (days_from_civil(year, month, day) - baseline) as f64; + serial += f64::from(hour * 3_600 + minute * 60 + second) / 86_400.0; + serial += f64::from(millis) / 86_400_000.0; + if (year, month, day) < (1900, 3, 1) && serial >= 2.0 { + serial -= 1.0; + } + if date_1904 { + serial -= 1_462.0; + } + Some(serial.to_string()) +} + +fn parse_date(value: &str) -> Option<(i32, u32, u32)> { + if value.len() != 10 || value.as_bytes()[4] != b'-' || value.as_bytes()[7] != b'-' { + return None; + } + let year = value[0..4].parse::().ok()?; + let month = value[5..7].parse::().ok()?; + let day = value[8..10].parse::().ok()?; + ((1..=12).contains(&month) && day >= 1 && day <= days_in_month(year, month)) + .then_some((year, month, day)) +} + +fn parse_time(value: &str) -> Option<(u32, u32, u32, u32)> { + let value = value.strip_suffix('Z').unwrap_or(value); + let (clock, millis) = value.split_once('.').map_or((value, 0), |(clock, millis)| { + if millis.len() != 3 { + return (clock, u32::MAX); + } + (clock, millis.parse::().unwrap_or(u32::MAX)) + }); + if clock.len() != 8 || clock.as_bytes()[2] != b':' || clock.as_bytes()[5] != b':' { + return None; + } + let hour = clock[0..2].parse::().ok()?; + let minute = clock[3..5].parse::().ok()?; + let second = clock[6..8].parse::().ok()?; + (hour < 24 && minute < 60 && second < 60 && millis < 1_000) + .then_some((hour, minute, second, millis)) +} + +fn days_in_month(year: i32, month: u32) -> u32 { + match month { + 1 | 3 | 5 | 7 | 8 | 10 | 12 => 31, + 4 | 6 | 9 | 11 => 30, + 2 if year % 400 == 0 || (year % 4 == 0 && year % 100 != 0) => 29, + 2 => 28, + _ => 0, + } +} + +fn days_from_civil(year: i32, month: u32, day: u32) -> i64 { + let year = year - i32::from(month <= 2); + let era = if year >= 0 { year } else { year - 399 } / 400; + let year_of_era = year - era * 400; + let month = i32::try_from(month).unwrap_or_default(); + let day = i32::try_from(day).unwrap_or_default(); + let day_of_year = (153 * (month + if month > 2 { -3 } else { 9 }) + 2) / 5 + day - 1; + let day_of_era = year_of_era * 365 + year_of_era / 4 - year_of_era / 100 + day_of_year; + i64::from(era * 146_097 + day_of_era) +} + +fn import_error(code: &str, message: impl Into) -> a3s_use_core::UseError { + super::super::editor_error(code, message) +} + +fn import_cell_count_limit() -> a3s_use_core::UseError { + import_error( + "use.office.spreadsheet_import_cell_count_limit", + format!( + "Spreadsheet import accepts at most {MAX_NATIVE_SPREADSHEET_IMPORT_CELLS} rectangular target cells." + ), + ) +} + +fn delimited_invalid( + rows: &[Vec], + row: &[String], + message: impl Into, +) -> a3s_use_core::UseError { + import_error("use.office.spreadsheet_import_delimited_invalid", message) + .with_detail("row", rows.len() + 1) + .with_detail("column", row.len() + 1) +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn parser_matches_delimited_quote_and_blank_line_semantics() { + assert_eq!( + parse_delimited("a,\"b,c\"\r\n\r\n\"d\nq\",\"x\"\"y\"", ',').unwrap(), + [ + vec!["a".to_string(), "b,c".to_string()], + vec![String::new()], + vec!["d\nq".to_string(), "x\"y".to_string()] + ] + ); + assert!(parse_delimited("", ',').unwrap().is_empty()); + assert_eq!(parse_delimited("\"\"", ',').unwrap(), [vec![String::new()]]); + } + + #[test] + fn parser_rejects_ambiguous_or_unclosed_quotes() { + for input in ["\"unclosed", "\"closed\"suffix", "unquoted\"quote"] { + let error = parse_delimited(input, ',').unwrap_err(); + assert_eq!( + error.code, + "use.office.spreadsheet_import_delimited_invalid" + ); + assert_eq!(error.details["row"], 1); + assert_eq!(error.details["column"], 1); + } + } + + #[test] + fn iso_dates_follow_the_declared_excel_date_system() { + let standard = parse_iso_date("2026-07-17T12:30:15.250Z", false) + .unwrap() + .parse::() + .unwrap(); + let date_1904 = parse_iso_date("2026-07-17T12:30:15.250Z", true) + .unwrap() + .parse::() + .unwrap(); + assert_eq!(standard - date_1904, 1_462.0); + assert_eq!(parse_iso_date("1900-02-28", false).as_deref(), Some("59")); + assert_eq!(parse_iso_date("1900-03-01", false).as_deref(), Some("61")); + assert!(parse_iso_date("2026-02-29", false).is_none()); + } +} diff --git a/crates/office/src/editor/spreadsheet/style.rs b/crates/office/src/editor/spreadsheet/style.rs index 64773242..26bf1807 100644 --- a/crates/office/src/editor/spreadsheet/style.rs +++ b/crates/office/src/editor/spreadsheet/style.rs @@ -199,6 +199,24 @@ fn set_format( package.set_part(&part_name, edited) } +pub(super) fn derived_cell_style_indexes( + package: &mut NativeOfficePackage, + base_styles: &BTreeSet, + format: &NativeSpreadsheetCellFormat, +) -> UseResult> { + format.validate()?; + ensure_style_collections(package)?; + let resolved = cell_format::resolve(package, Some(format))?; + base_styles + .iter() + .copied() + .map(|base_style| { + style_index_for_format(package, base_style, None, Some(format), &resolved) + .map(|derived| (base_style, derived)) + }) + .collect() +} + fn ensure_style_collections(package: &mut NativeOfficePackage) -> UseResult<()> { ensure_styles_part(package)?; ensure_collection(package, "fonts", "font", STYLE_CHILDREN_AFTER_FONTS)?; @@ -686,7 +704,7 @@ fn styled_cell(prefix: Option<&str>, reference: &str, style: usize) -> String { format!("<{tag} r=\"{reference}\" s=\"{style}\"/>") } -fn cell_style_index(cell: &IndexedXmlElement) -> UseResult { +pub(super) fn cell_style_index(cell: &IndexedXmlElement) -> UseResult { cell.attributes.get("s").map_or(Ok(0), |value| { value.parse::().map_err(|_| styles_invalid()) }) diff --git a/crates/office/src/editor/spreadsheet/table.rs b/crates/office/src/editor/spreadsheet/table.rs index 58bd7b14..b5852851 100644 --- a/crates/office/src/editor/spreadsheet/table.rs +++ b/crates/office/src/editor/spreadsheet/table.rs @@ -1,13 +1,17 @@ +use std::collections::BTreeMap; + use a3s_use_core::UseResult; use super::{editor_error, validate_mutation_path, SpreadsheetCellValue}; use crate::editor::part::{dialect, relationship_part, relative_target}; use crate::editor::NativeSpreadsheetTable; use crate::semantic::{DocumentNode, NativeOfficeDocument, OfficeNodeType}; +use crate::spreadsheet_formula::StructuredReferenceRewritePlan; use crate::spreadsheet_reference::{CellRange, CellReference}; use crate::xml_edit::index_xml; use crate::{DocumentKind, NativeOfficePackage, RelationshipSource, RelationshipTarget}; +mod formula; mod xml; const MAX_SPREADSHEET_TABLES: usize = 65_536; @@ -97,7 +101,7 @@ pub(super) fn set( let range = table.validate()?; table.range = range.a1(); let snapshot = NativeOfficeDocument::from_package(package.clone())?; - let node = snapshot.get(&resolved.path, 0)?; + let node = snapshot.get(&resolved.path, 3)?; if node.format.get("nativeMutable").map(String::as_str) != Some("true") { return Err(editor_error( "use.office.spreadsheet_table_unknown_content", @@ -110,6 +114,8 @@ pub(super) fn set( "Keep the imported table unchanged or inspect its OOXML before replacing it through the typed table contract.", )); } + let old_table = NativeSpreadsheetTable::from_semantic_node(&node)?; + let old_range = CellRange::parse(&old_table.range)?; validate_identity(&snapshot, &table, Some(&resolved.path))?; validate_range( package, @@ -118,35 +124,118 @@ pub(super) fn set( range, Some(&resolved.path), )?; - let part = package.xml_part(&resolved.part)?; + let (plan, formula_rewrite_required) = table_formula_rewrite_plan( + &old_table, + &table, + old_range, + range, + resolved.sheet.path.trim_start_matches('/'), + ); + let mut candidate = package.clone(); + if formula_rewrite_required { + formula::rewrite_table_references( + &mut candidate, + &resolved.sheet.part, + &resolved.part, + old_range, + &plan, + )?; + } + let part = candidate.xml_part(&resolved.part)?; let edited = xml::replace_table(&part, &table, range)?; - package.set_part(&resolved.part, edited)?; - stamp_headers(package, &resolved.sheet, &table, range)?; - super::mark_workbook_for_recalculation(package)?; + candidate.set_part(&resolved.part, edited)?; + stamp_headers(&mut candidate, &resolved.sheet, &table, range)?; + super::mark_workbook_for_recalculation(&mut candidate)?; + *package = candidate; Ok(resolved.path) } pub(super) fn remove(package: &mut NativeOfficePackage, path: &str) -> UseResult<()> { let resolved = resolve_table(package, path)?; validate_relationship_graph(package, &resolved)?; - let worksheet = package.xml_part(&resolved.sheet.part)?; + let snapshot = NativeOfficeDocument::from_package(package.clone())?; + let node = snapshot.get(&resolved.path, 0)?; + let name = required_table_format(&node, "name")?.to_string(); + let display_name = required_table_format(&node, "displayName")?.to_string(); + let range = CellRange::parse(required_table_format(&node, "ref")?)?; + let plan = StructuredReferenceRewritePlan::removal( + name.clone(), + resolved.sheet.path.trim_start_matches('/'), + [name, display_name], + ); + let mut candidate = package.clone(); + formula::rewrite_table_references( + &mut candidate, + &resolved.sheet.part, + &resolved.part, + range, + &plan, + )?; + let worksheet = candidate.xml_part(&resolved.sheet.part)?; let worksheet = xml::remove_table_part_reference(&worksheet, &resolved.relationship_id)?; - package.set_part(&resolved.sheet.part, worksheet)?; + candidate.set_part(&resolved.sheet.part, worksheet)?; crate::opc_edit::remove_relationship( - package, + &mut candidate, &relationship_part(&resolved.sheet.part), &resolved.relationship_id, )?; - let content_types = package.opc_model()?.content_types().clone(); + let content_types = candidate.opc_model()?.content_types().clone(); if content_types.override_for_part(&resolved.part).is_some() { - crate::opc_edit::remove_content_type_override(package, &resolved.part)?; + crate::opc_edit::remove_content_type_override(&mut candidate, &resolved.part)?; } - package.remove_part(&resolved.part)?; + candidate.remove_part(&resolved.part)?; let table_relationships = relationship_part(&resolved.part); - if package.contains_part(&table_relationships) { - package.remove_part(&table_relationships)?; + if candidate.contains_part(&table_relationships) { + candidate.remove_part(&table_relationships)?; + } + super::mark_workbook_for_recalculation(&mut candidate)?; + *package = candidate; + Ok(()) +} + +fn table_formula_rewrite_plan( + old: &NativeSpreadsheetTable, + new: &NativeSpreadsheetTable, + old_range: CellRange, + new_range: CellRange, + sheet: &str, +) -> (StructuredReferenceRewritePlan, bool) { + let old_display_name = old.display_name.as_deref().unwrap_or(&old.name); + let new_display_name = new.display_name.as_deref().unwrap_or(&new.name); + let mut aliases = BTreeMap::new(); + aliases.insert(old.name.to_lowercase(), new.name.clone()); + aliases.insert( + old_display_name.to_lowercase(), + new_display_name.to_string(), + ); + + let mut columns = BTreeMap::new(); + for (index, old_column) in old.columns.iter().enumerate() { + let replacement = new.columns.get(index).map(|column| column.name.clone()); + if replacement.as_deref() != Some(old_column.name.as_str()) { + columns.insert(old_column.name.to_lowercase(), replacement); + } } - super::mark_workbook_for_recalculation(package) + + let aliases_changed = if old.name.eq_ignore_ascii_case(old_display_name) { + old_display_name != new_display_name + } else { + old.name != new.name || old_display_name != new_display_name + }; + let geometry_changed = old_range != new_range + || old.header_row != new.header_row + || old.totals_row != new.totals_row; + let rewrite_required = aliases_changed || !columns.is_empty() || geometry_changed; + ( + StructuredReferenceRewritePlan::rename( + old.name.clone(), + sheet, + aliases, + columns, + geometry_changed, + ), + rewrite_required, + ) } fn resolve_sheet(package: &NativeOfficePackage, requested: &str) -> UseResult { @@ -547,6 +636,16 @@ fn spreadsheet_tables(document: &NativeOfficeDocument) -> impl Iterator(node: &'a DocumentNode, key: &str) -> UseResult<&'a str> { + node.format.get(key).map(String::as_str).ok_or_else(|| { + editor_error( + "use.office.spreadsheet_table_invalid", + format!("Spreadsheet table '{}' has no '{key}' property.", node.path), + ) + .with_detail("path", node.path.clone()) + }) +} + fn identity_collision(name: &str, owner: &str) -> a3s_use_core::UseError { editor_error( "use.office.spreadsheet_table_name_collision", diff --git a/crates/office/src/editor/spreadsheet/table/formula.rs b/crates/office/src/editor/spreadsheet/table/formula.rs new file mode 100644 index 00000000..71f21301 --- /dev/null +++ b/crates/office/src/editor/spreadsheet/table/formula.rs @@ -0,0 +1,299 @@ +use a3s_use_core::UseResult; + +use crate::semantic::{NativeOfficeDocument, OfficeNodeType}; +use crate::spreadsheet_formula::{ + LocalStructuredReferenceContext, StructuredReferenceRewritePlan, + StructuredReferenceRewriteResult, +}; +use crate::spreadsheet_reference::{CellRange, CellReference}; +use crate::xml_edit::{apply_patches, escape_text, index_xml, IndexedXmlElement, XmlPatch}; +use crate::{LosslessXmlPart, NativeOfficePackage}; + +#[derive(Debug, Clone, Copy)] +enum FormulaPartKind { + Workbook { + target_sheet_index: usize, + }, + Worksheet { + target: bool, + table_range: CellRange, + }, + Chart, + Table { + target: bool, + }, +} + +#[derive(Debug, Clone, Copy)] +struct FormulaCarrier<'a> { + element: &'a IndexedXmlElement, + context: LocalStructuredReferenceContext, +} + +pub(super) fn rewrite_table_references( + package: &mut NativeOfficePackage, + target_sheet_part: &str, + target_table_part: &str, + table_range: CellRange, + plan: &StructuredReferenceRewritePlan, +) -> UseResult<()> { + let snapshot = NativeOfficeDocument::from_package(package.clone())?; + let worksheet_parts = snapshot + .root() + .children + .iter() + .filter(|node| node.node_type == OfficeNodeType::Worksheet) + .map(|node| { + node.format.get("part").cloned().ok_or_else(|| { + formula_error( + "use.office.spreadsheet_sheet_invalid", + format!("Worksheet '{}' has no source part.", node.path), + ) + }) + }) + .collect::>>()?; + let target_sheet_index = worksheet_parts + .iter() + .position(|part| part == target_sheet_part) + .ok_or_else(|| { + formula_error( + "use.office.spreadsheet_sheet_invalid", + format!("Worksheet part '{target_sheet_part}' has no workbook sheet entry."), + ) + })?; + let chart_parts = package + .part_names() + .filter(|part| part.starts_with("xl/charts/") && part.ends_with(".xml")) + .map(str::to_string) + .collect::>(); + let table_parts = package + .part_names() + .filter(|part| part.starts_with("xl/tables/") && part.ends_with(".xml")) + .map(str::to_string) + .collect::>(); + + let mut matched = rewrite_part( + package, + "xl/workbook.xml", + FormulaPartKind::Workbook { target_sheet_index }, + plan, + )?; + for part_name in &worksheet_parts { + matched |= rewrite_part( + package, + part_name, + FormulaPartKind::Worksheet { + target: part_name == target_sheet_part, + table_range, + }, + plan, + )?; + } + for part_name in &chart_parts { + matched |= rewrite_part(package, part_name, FormulaPartKind::Chart, plan)?; + } + for part_name in &table_parts { + matched |= rewrite_part( + package, + part_name, + FormulaPartKind::Table { + target: part_name == target_table_part, + }, + plan, + )?; + } + + if matched && plan.geometry_changed() { + clear_formula_caches(package, &worksheet_parts)?; + clear_chart_caches(package, &chart_parts)?; + } + Ok(()) +} + +fn rewrite_part( + package: &mut NativeOfficePackage, + part_name: &str, + kind: FormulaPartKind, + plan: &StructuredReferenceRewritePlan, +) -> UseResult { + let part = package.xml_part(part_name)?; + let root = index_xml(&part)?; + let mut carriers = Vec::new(); + collect_formula_carriers(&root, None, kind, &mut carriers)?; + let mut patches = Vec::new(); + let mut matched = false; + for carrier in carriers { + let formula = decoded_text(&part, carrier.element)?; + let rewritten: StructuredReferenceRewriteResult = + plan.rewrite(&formula, carrier.context)?; + matched |= rewritten.matched; + if rewritten.formula != formula { + patches.push(XmlPatch::new( + carrier.element.content_range.clone(), + escape_text(&rewritten.formula), + )); + } + } + if !patches.is_empty() { + package.set_part(part_name, apply_patches(&part, patches)?)?; + } + Ok(matched) +} + +fn collect_formula_carriers<'a>( + element: &'a IndexedXmlElement, + parent: Option<&'a IndexedXmlElement>, + kind: FormulaPartKind, + output: &mut Vec>, +) -> UseResult<()> { + if is_formula_element(element, kind) { + output.push(FormulaCarrier { + element, + context: local_context(element, parent, kind)?, + }); + } + for child in &element.children { + collect_formula_carriers(child, Some(element), kind, output)?; + } + Ok(()) +} + +fn is_formula_element(element: &IndexedXmlElement, kind: FormulaPartKind) -> bool { + if matches!(kind, FormulaPartKind::Workbook { .. }) { + return element.local_name == "definedName"; + } + matches!( + element.local_name.as_str(), + "f" | "formula" | "formula1" | "formula2" | "calculatedColumnFormula" | "totalsRowFormula" + ) +} + +fn local_context( + element: &IndexedXmlElement, + parent: Option<&IndexedXmlElement>, + kind: FormulaPartKind, +) -> UseResult { + match kind { + FormulaPartKind::Workbook { target_sheet_index } => { + let Some(local_sheet_id) = element.attributes.get("localSheetId") else { + return Ok(LocalStructuredReferenceContext::Unknown); + }; + Ok(local_sheet_id + .parse::() + .ok() + .filter(|index| *index == target_sheet_index) + .map_or(LocalStructuredReferenceContext::DoesNotApply, |_| { + LocalStructuredReferenceContext::Unknown + })) + } + FormulaPartKind::Worksheet { + target, + table_range, + } => { + if !target { + return Ok(LocalStructuredReferenceContext::DoesNotApply); + } + if element.local_name != "f" + || parent.map(|value| value.local_name.as_str()) != Some("c") + { + return Ok(LocalStructuredReferenceContext::Unknown); + } + let reference = parent + .and_then(|cell| cell.attributes.get("r")) + .ok_or_else(|| { + formula_error( + "use.office.spreadsheet_formula_invalid", + "Spreadsheet formula cell has no cell reference.", + ) + }) + .and_then(|reference| CellReference::parse(reference))?; + Ok(if table_range.contains(reference) { + LocalStructuredReferenceContext::Applies + } else { + LocalStructuredReferenceContext::DoesNotApply + }) + } + FormulaPartKind::Chart => Ok(LocalStructuredReferenceContext::Unknown), + FormulaPartKind::Table { target } => Ok(if target { + LocalStructuredReferenceContext::Applies + } else { + LocalStructuredReferenceContext::DoesNotApply + }), + } +} + +fn clear_formula_caches( + package: &mut NativeOfficePackage, + worksheet_parts: &[String], +) -> UseResult<()> { + for part_name in worksheet_parts { + let part = package.xml_part(part_name)?; + let root = index_xml(&part)?; + let mut cells = Vec::new(); + root.descendants_named("c", &mut cells); + let patches = cells + .into_iter() + .filter(|cell| cell.children.iter().any(|child| child.local_name == "f")) + .filter_map(|cell| cell.children.iter().find(|child| child.local_name == "v")) + .map(|value| XmlPatch::new(value.full_range.clone(), Vec::new())) + .collect::>(); + if !patches.is_empty() { + package.set_part(part_name, apply_patches(&part, patches)?)?; + } + } + Ok(()) +} + +fn clear_chart_caches(package: &mut NativeOfficePackage, chart_parts: &[String]) -> UseResult<()> { + for part_name in chart_parts { + let part = package.xml_part(part_name)?; + let root = index_xml(&part)?; + let mut patches = Vec::new(); + for cache_name in ["numCache", "strCache", "multiLvlStrCache"] { + let mut caches = Vec::new(); + root.descendants_named(cache_name, &mut caches); + patches.extend( + caches + .into_iter() + .map(|cache| XmlPatch::new(cache.full_range.clone(), Vec::new())), + ); + } + if !patches.is_empty() { + package.set_part(part_name, apply_patches(&part, patches)?)?; + } + } + Ok(()) +} + +fn decoded_text(part: &LosslessXmlPart, element: &IndexedXmlElement) -> UseResult { + let bytes = part + .parse_bytes() + .get(element.content_range.clone()) + .ok_or_else(|| { + formula_error( + "use.office.spreadsheet_formula_invalid", + format!("Formula range in '{}' is invalid.", part.name()), + ) + })?; + let text = std::str::from_utf8(bytes).map_err(|error| { + formula_error( + "use.office.spreadsheet_formula_invalid", + format!("Formula in '{}' is not UTF-8: {error}", part.name()), + ) + })?; + quick_xml::escape::unescape(text) + .map(|value| value.into_owned()) + .map_err(|error| { + formula_error( + "use.office.spreadsheet_formula_invalid", + format!( + "Formula in '{}' contains invalid XML escapes: {error}", + part.name() + ), + ) + }) +} + +fn formula_error(code: &str, message: impl Into) -> a3s_use_core::UseError { + super::super::editor_error(code, message) +} diff --git a/crates/office/src/editor/spreadsheet/view.rs b/crates/office/src/editor/spreadsheet/view.rs new file mode 100644 index 00000000..20280bba --- /dev/null +++ b/crates/office/src/editor/spreadsheet/view.rs @@ -0,0 +1,300 @@ +use a3s_use_core::UseResult; + +use super::{editor_error, prefix, qualified, validate_mutation_path}; +use crate::xml_edit::{apply_patches, index_xml, insert_child, insert_ordered_child, XmlPatch}; +use crate::{ + DocumentKind, NativeOfficeDocument, NativeOfficePackage, NativeSpreadsheetFrozenPane, + OfficeNodeType, +}; + +const WORKSHEET_CHILDREN_AFTER_VIEWS: &[&str] = &[ + "sheetFormatPr", + "cols", + "sheetData", + "sheetCalcPr", + "sheetProtection", + "protectedRanges", + "scenarios", + "autoFilter", + "sortState", + "dataConsolidate", + "customSheetViews", + "mergeCells", + "phoneticPr", + "conditionalFormatting", + "dataValidations", + "hyperlinks", + "printOptions", + "pageMargins", + "pageSetup", + "headerFooter", + "rowBreaks", + "colBreaks", + "customProperties", + "cellWatches", + "ignoredErrors", + "smartTags", + "drawing", + "legacyDrawing", + "legacyDrawingHF", + "picture", + "oleObjects", + "controls", + "webPublishItems", + "tableParts", + "extLst", +]; + +struct ResolvedSheet { + path: String, + part: String, +} + +pub(super) fn is_path(path: &str) -> bool { + let normalized = path.trim_matches('/'); + normalized + .rsplit_once('/') + .is_some_and(|(parent, segment)| { + !parent.contains('/') && segment.eq_ignore_ascii_case("freeze") + }) +} + +pub(super) fn set( + package: &mut NativeOfficePackage, + requested: &str, + pane: &NativeSpreadsheetFrozenPane, +) -> UseResult { + let sheet = resolve_sheet(package, requested)?; + let pane = pane.normalized()?; + let worksheet = package.xml_part(&sheet.part)?; + let root = index_xml(&worksheet)?; + let views = root + .children + .iter() + .filter(|child| child.local_name == "sheetViews" && child.namespace == root.namespace) + .collect::>(); + if views.len() > 1 { + return Err(view_part_error( + &sheet.part, + "contains multiple sheetViews collections", + )); + } + let fragment = pane_fragment(prefix(&root.qualified_name), &pane); + let edited = if let Some(views) = views.first().copied() { + let matching = views + .children + .iter() + .filter(|child| { + child.local_name == "sheetView" + && child.namespace == root.namespace + && child.attributes.get("workbookViewId").map(String::as_str) == Some("0") + }) + .collect::>(); + if matching.len() > 1 { + return Err(view_part_error( + &sheet.part, + "contains multiple sheetView elements for workbookViewId 0", + )); + } + if let Some(view) = matching.first().copied() { + let panes = direct_panes(view, root.namespace.as_deref()); + if panes.len() > 1 { + return Err(view_part_error( + &sheet.part, + "contains multiple pane elements in workbookViewId 0", + )); + } + if let Some(existing) = panes.first().copied() { + require_mutable_pane(&worksheet, existing)?; + apply_patches( + &worksheet, + vec![XmlPatch::new(existing.full_range.clone(), fragment)], + )? + } else { + insert_ordered_child( + &worksheet, + view, + fragment, + &["selection", "pivotSelection", "extLst"], + )? + } + } else { + let tag = qualified(prefix(&views.qualified_name), "sheetView"); + insert_child( + &worksheet, + views, + format!("<{tag} workbookViewId=\"0\">{fragment}"), + )? + } + } else { + let view_prefix = prefix(&root.qualified_name); + let views_tag = qualified(view_prefix, "sheetViews"); + let view_tag = qualified(view_prefix, "sheetView"); + insert_ordered_child( + &worksheet, + &root, + format!( + "<{views_tag}><{view_tag} workbookViewId=\"0\">{fragment}" + ), + WORKSHEET_CHILDREN_AFTER_VIEWS, + )? + }; + package.set_part(&sheet.part, edited)?; + Ok(format!("{}/freeze", sheet.path)) +} + +pub(super) fn remove(package: &mut NativeOfficePackage, requested: &str) -> UseResult<()> { + validate_mutation_path(requested)?; + if !is_path(requested) { + return Err(editor_error( + "use.office.mutation_path_unsupported", + "Removing a frozen Spreadsheet pane requires a path such as /Sheet1/freeze.", + )); + } + let snapshot = NativeOfficeDocument::from_package(package.clone())?; + let node = snapshot.get(requested, 0)?; + if node.node_type != OfficeNodeType::FrozenPane { + return Err(node_not_found(requested)); + } + if node.format.get("nativeMutable").map(String::as_str) != Some("true") { + return Err(editor_error( + "use.office.spreadsheet_freeze_unknown_content", + format!( + "Frozen Spreadsheet pane '{}' contains unsupported view state or unknown content.", + node.path + ), + )); + } + let (sheet_path, _) = node + .path + .rsplit_once('/') + .ok_or_else(|| node_not_found(requested))?; + let sheet = resolve_sheet(package, sheet_path)?; + let worksheet = package.xml_part(&sheet.part)?; + let root = index_xml(&worksheet)?; + let panes = root + .children + .iter() + .filter(|child| child.local_name == "sheetViews" && child.namespace == root.namespace) + .flat_map(|views| views.children.iter()) + .filter(|view| { + view.local_name == "sheetView" + && view.namespace == root.namespace + && view.attributes.get("workbookViewId").map(String::as_str) == Some("0") + }) + .flat_map(|view| direct_panes(view, root.namespace.as_deref())) + .collect::>(); + if panes.len() != 1 { + return Err(view_part_error( + &sheet.part, + "does not contain exactly one frozen pane for workbookViewId 0", + )); + } + require_mutable_pane(&worksheet, panes[0])?; + let edited = apply_patches( + &worksheet, + vec![XmlPatch::new(panes[0].full_range.clone(), Vec::new())], + )?; + package.set_part(&sheet.part, edited) +} + +fn resolve_sheet(package: &NativeOfficePackage, requested: &str) -> UseResult { + if package.kind() != DocumentKind::Spreadsheet { + return Err(editor_error( + "use.office.mutation_type_unsupported", + "Frozen pane operations are available only for Spreadsheet documents.", + )); + } + validate_mutation_path(requested)?; + if requested.trim_start_matches('/').contains('/') { + return Err(editor_error( + "use.office.mutation_path_unsupported", + "Setting a frozen Spreadsheet pane requires a worksheet path such as /Sheet1.", + )); + } + let snapshot = NativeOfficeDocument::from_package(package.clone())?; + let sheet = snapshot + .root() + .children + .iter() + .find(|node| { + node.node_type == OfficeNodeType::Worksheet && node.path.eq_ignore_ascii_case(requested) + }) + .ok_or_else(|| node_not_found(requested))?; + let part = sheet.format.get("part").cloned().ok_or_else(|| { + editor_error( + "use.office.spreadsheet_sheet_invalid", + format!("Worksheet '{}' has no source part.", sheet.path), + ) + })?; + Ok(ResolvedSheet { + path: sheet.path.clone(), + part, + }) +} + +fn pane_fragment(prefix: Option<&str>, pane: &NativeSpreadsheetFrozenPane) -> String { + let tag = qualified(prefix, "pane"); + let x_split = if pane.frozen_columns > 0 { + format!(" xSplit=\"{}\"", pane.frozen_columns) + } else { + String::new() + }; + let y_split = if pane.frozen_rows > 0 { + format!(" ySplit=\"{}\"", pane.frozen_rows) + } else { + String::new() + }; + format!( + "<{tag}{x_split}{y_split} topLeftCell=\"{}\" activePane=\"{}\" state=\"frozen\"/>", + pane.top_left_cell, + pane.active_pane() + ) +} + +fn direct_panes<'a>( + view: &'a crate::xml_edit::IndexedXmlElement, + namespace: Option<&str>, +) -> Vec<&'a crate::xml_edit::IndexedXmlElement> { + view.children + .iter() + .filter(|child| child.local_name == "pane" && child.namespace.as_deref() == namespace) + .collect() +} + +fn require_mutable_pane( + part: &crate::LosslessXmlPart, + pane: &crate::xml_edit::IndexedXmlElement, +) -> UseResult<()> { + let known = ["xSplit", "ySplit", "topLeftCell", "activePane", "state"]; + let unknown_attribute = pane + .qualified_attributes + .keys() + .any(|name| !known.contains(&name.as_str())); + let content = std::str::from_utf8(&part.parse_bytes()[pane.content_range.clone()]) + .unwrap_or("") + .trim(); + if unknown_attribute || !pane.children.is_empty() || !content.is_empty() { + return Err(editor_error( + "use.office.spreadsheet_freeze_unknown_content", + "Frozen Spreadsheet pane contains unknown attributes or child content.", + )); + } + Ok(()) +} + +fn node_not_found(path: &str) -> a3s_use_core::UseError { + editor_error( + "use.office.node_not_found", + format!("Office semantic path '{path}' does not exist."), + ) + .with_detail("path", path) +} + +fn view_part_error(part: &str, reason: &str) -> a3s_use_core::UseError { + editor_error( + "use.office.spreadsheet_freeze_invalid", + format!("Spreadsheet worksheet part '{part}' {reason}."), + ) + .with_detail("part", part) +} diff --git a/crates/office/src/editor/types.rs b/crates/office/src/editor/types.rs index 6a577ce5..0b0c9440 100644 --- a/crates/office/src/editor/types.rs +++ b/crates/office/src/editor/types.rs @@ -4,14 +4,17 @@ use serde::{Deserialize, Serialize}; use url::Url; use super::part::{NativeCreatedPart, NativeOfficePartType}; +use crate::spreadsheet_formula::SpreadsheetFormulaCalculation; mod conditional_formatting; mod data_validation; mod formatting; mod named_range; mod spreadsheet_filter; +mod spreadsheet_import; mod spreadsheet_sort; mod spreadsheet_table; +mod spreadsheet_view; pub use conditional_formatting::{ NativeSpreadsheetConditionalFormat, NativeSpreadsheetConditionalFormatIconSet, @@ -35,12 +38,18 @@ pub use spreadsheet_filter::{ NativeSpreadsheetAutoFilter, NativeSpreadsheetDynamicFilter, NativeSpreadsheetFilterColumn, NativeSpreadsheetFilterCriteria, }; +pub use spreadsheet_import::{ + NativeSpreadsheetDelimitedFormat, NativeSpreadsheetDelimitedImport, + NativeSpreadsheetImportResult, MAX_NATIVE_SPREADSHEET_IMPORT_BYTES, + MAX_NATIVE_SPREADSHEET_IMPORT_CELLS, +}; pub use spreadsheet_sort::{ NativeSpreadsheetSort, NativeSpreadsheetSortDirection, NativeSpreadsheetSortKey, }; pub use spreadsheet_table::{ NativeSpreadsheetTable, NativeSpreadsheetTableColumn, NativeSpreadsheetTableStyle, }; +pub use spreadsheet_view::NativeSpreadsheetFrozenPane; /// A hyperlink destination represented without executing or resolving it. #[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] @@ -677,6 +686,7 @@ pub enum NativeOfficeMutation { path: String, value: SpreadsheetCellValue, }, + RecalculateSpreadsheetFormulas, AddSpreadsheetTable { sheet: String, table: NativeSpreadsheetTable, @@ -697,6 +707,14 @@ pub enum NativeOfficeMutation { path: String, sort: NativeSpreadsheetSort, }, + ImportSpreadsheetDelimited { + sheet: String, + import: NativeSpreadsheetDelimitedImport, + }, + SetSpreadsheetFrozenPane { + sheet: String, + pane: NativeSpreadsheetFrozenPane, + }, AddNamedRange { #[serde(rename = "namedRange")] named_range: NativeSpreadsheetNamedRange, @@ -862,4 +880,8 @@ pub struct NativeBatchResult { pub created_images: Vec, #[serde(default, skip_serializing_if = "Vec::is_empty")] pub text_replacements: Vec, + #[serde(default, skip_serializing_if = "Vec::is_empty")] + pub spreadsheet_imports: Vec, + #[serde(default, skip_serializing_if = "Vec::is_empty")] + pub spreadsheet_calculations: Vec, } diff --git a/crates/office/src/editor/types/spreadsheet_import.rs b/crates/office/src/editor/types/spreadsheet_import.rs new file mode 100644 index 00000000..fa879c6a --- /dev/null +++ b/crates/office/src/editor/types/spreadsheet_import.rs @@ -0,0 +1,109 @@ +use a3s_use_core::UseResult; +use serde::{Deserialize, Serialize}; + +use crate::spreadsheet_reference::CellReference; + +/// Maximum UTF-8 source bytes accepted by one native delimited import. +pub const MAX_NATIVE_SPREADSHEET_IMPORT_BYTES: usize = 8 * 1024 * 1024; +/// Maximum rectangular target cells admitted by one native delimited import. +pub const MAX_NATIVE_SPREADSHEET_IMPORT_CELLS: usize = 100_000; + +/// Closed source syntax accepted by native Spreadsheet import. +#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)] +#[serde(rename_all = "kebab-case")] +pub enum NativeSpreadsheetDelimitedFormat { + Csv, + Tsv, +} + +impl NativeSpreadsheetDelimitedFormat { + pub(crate) fn delimiter(self) -> char { + match self { + Self::Csv => ',', + Self::Tsv => '\t', + } + } +} + +/// Bounded, filesystem-independent source for one Spreadsheet import. +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +#[serde(rename_all = "camelCase", deny_unknown_fields)] +pub struct NativeSpreadsheetDelimitedImport { + pub content: String, + pub format: NativeSpreadsheetDelimitedFormat, + #[serde(default)] + pub header: bool, + #[serde(default = "default_start_cell")] + pub start_cell: String, +} + +impl NativeSpreadsheetDelimitedImport { + pub fn new(content: impl Into, format: NativeSpreadsheetDelimitedFormat) -> Self { + Self { + content: content.into(), + format, + header: false, + start_cell: default_start_cell(), + } + } + + pub fn with_header(mut self, header: bool) -> Self { + self.header = header; + self + } + + pub fn with_start_cell(mut self, start_cell: impl Into) -> Self { + self.start_cell = start_cell.into(); + self + } + + pub(crate) fn validate(&self) -> UseResult { + if self.content.len() > MAX_NATIVE_SPREADSHEET_IMPORT_BYTES { + return Err(import_error( + "use.office.spreadsheet_import_input_limit", + format!( + "Native Spreadsheet delimited import accepts at most {MAX_NATIVE_SPREADSHEET_IMPORT_BYTES} UTF-8 bytes." + ), + ) + .with_detail("bytes", self.content.len())); + } + CellReference::parse(&self.start_cell).map_err(|error| { + import_error( + "use.office.spreadsheet_import_start_cell_invalid", + format!( + "Spreadsheet import startCell '{}' is invalid: {error}", + self.start_cell + ), + ) + .with_detail("startCell", self.start_cell.clone()) + }) + } +} + +/// Receipt returned for one atomic native delimited import mutation. +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +#[serde(rename_all = "camelCase")] +pub struct NativeSpreadsheetImportResult { + pub path: String, + pub sheet: String, + pub start_cell: String, + #[serde(skip_serializing_if = "Option::is_none")] + pub range: Option, + pub format: NativeSpreadsheetDelimitedFormat, + pub row_count: usize, + pub column_count: usize, + pub header: bool, + pub changed: bool, + #[serde(skip_serializing_if = "Option::is_none")] + pub filter_path: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub freeze_path: Option, +} + +fn default_start_cell() -> String { + "A1".to_string() +} + +fn import_error(code: &str, message: impl Into) -> a3s_use_core::UseError { + super::super::editor_error(code, message) +} diff --git a/crates/office/src/editor/types/spreadsheet_view.rs b/crates/office/src/editor/types/spreadsheet_view.rs new file mode 100644 index 00000000..b3894500 --- /dev/null +++ b/crates/office/src/editor/types/spreadsheet_view.rs @@ -0,0 +1,130 @@ +use a3s_use_core::UseResult; +use serde::{Deserialize, Serialize}; + +use crate::semantic::{DocumentNode, OfficeNodeType}; +use crate::spreadsheet_reference::CellReference; + +const MAX_ROWS: u32 = 1_048_576; +const MAX_COLUMNS: u32 = 16_384; + +/// One canonical frozen Spreadsheet pane. +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +#[serde(rename_all = "camelCase", deny_unknown_fields)] +pub struct NativeSpreadsheetFrozenPane { + pub frozen_rows: u32, + pub frozen_columns: u32, + pub top_left_cell: String, +} + +impl NativeSpreadsheetFrozenPane { + pub fn new(frozen_rows: u32, frozen_columns: u32, top_left_cell: impl Into) -> Self { + Self { + frozen_rows, + frozen_columns, + top_left_cell: top_left_cell.into(), + } + } + + pub(crate) fn normalized(&self) -> UseResult { + if self.frozen_rows == 0 && self.frozen_columns == 0 { + return Err(view_error( + "use.office.spreadsheet_freeze_empty", + "A frozen Spreadsheet pane requires at least one frozen row or column.", + )); + } + if self.frozen_rows >= MAX_ROWS { + return Err(view_error( + "use.office.spreadsheet_freeze_row_limit", + format!("Frozen Spreadsheet rows must be below {MAX_ROWS}."), + )); + } + if self.frozen_columns >= MAX_COLUMNS { + return Err(view_error( + "use.office.spreadsheet_freeze_column_limit", + format!("Frozen Spreadsheet columns must be below {MAX_COLUMNS}."), + )); + } + let top_left = CellReference::parse(&self.top_left_cell).map_err(|error| { + view_error( + "use.office.spreadsheet_freeze_cell_invalid", + format!( + "Frozen Spreadsheet topLeftCell '{}' is invalid: {error}", + self.top_left_cell + ), + ) + })?; + if top_left.row <= self.frozen_rows || top_left.column <= self.frozen_columns { + return Err(view_error( + "use.office.spreadsheet_freeze_geometry_invalid", + "Frozen Spreadsheet topLeftCell must be below and to the right of every frozen split.", + )); + } + Ok(Self { + frozen_rows: self.frozen_rows, + frozen_columns: self.frozen_columns, + top_left_cell: top_left.a1(), + }) + } + + pub fn from_semantic_node(node: &DocumentNode) -> UseResult { + if node.node_type != OfficeNodeType::FrozenPane { + return Err(view_error( + "use.office.spreadsheet_freeze_node_invalid", + format!( + "Office node '{}' is not a frozen Spreadsheet pane.", + node.path + ), + )); + } + let parse = |key: &str| -> UseResult { + node.format + .get(key) + .ok_or_else(|| { + view_error( + "use.office.spreadsheet_freeze_node_invalid", + format!( + "Frozen Spreadsheet pane '{}' has no {key} value.", + node.path + ), + ) + })? + .parse::() + .map_err(|error| { + view_error( + "use.office.spreadsheet_freeze_node_invalid", + format!( + "Frozen Spreadsheet pane '{}' has invalid {key}: {error}", + node.path + ), + ) + }) + }; + Self::new( + parse("frozenRows")?, + parse("frozenColumns")?, + node.format.get("topLeftCell").cloned().ok_or_else(|| { + view_error( + "use.office.spreadsheet_freeze_node_invalid", + format!( + "Frozen Spreadsheet pane '{}' has no topLeftCell value.", + node.path + ), + ) + })?, + ) + .normalized() + } + + pub(crate) fn active_pane(&self) -> &'static str { + match (self.frozen_rows > 0, self.frozen_columns > 0) { + (true, true) => "bottomRight", + (true, false) => "bottomLeft", + (false, true) => "topRight", + (false, false) => "topLeft", + } + } +} + +fn view_error(code: &str, message: impl Into) -> a3s_use_core::UseError { + super::super::editor_error(code, message) +} diff --git a/crates/office/src/issues.rs b/crates/office/src/issues.rs index 51b91a6b..69a1ad91 100644 --- a/crates/office/src/issues.rs +++ b/crates/office/src/issues.rs @@ -194,12 +194,12 @@ impl<'a> IssueScanner<'a> { "Formula has an error result.", "Inspect the formula inputs and replace the invalid expression or references.", ) - } else if node.text.is_empty() { + } else if node.format.get("formulaCached").map(String::as_str) == Some("false") { ( NativeOfficeIssueSubtype::FormulaNotEvaluated, NativeOfficeIssueSeverity::Warning, "Formula has no cached result and requires recalculation.", - "Open the workbook in a conforming spreadsheet engine or run a future native recalculation pass.", + "Run `a3s use office native recalculate` or open the workbook in a conforming spreadsheet engine.", ) } else { return; diff --git a/crates/office/src/lib.rs b/crates/office/src/lib.rs index d519b5c5..ad09d5e1 100644 --- a/crates/office/src/lib.rs +++ b/crates/office/src/lib.rs @@ -56,15 +56,17 @@ pub use editor::{ NativeSpreadsheetConditionalFormatThreshold, NativeSpreadsheetConditionalFormatThresholdKind, NativeSpreadsheetConditionalFormatTimePeriod, NativeSpreadsheetDataValidation, NativeSpreadsheetDataValidationErrorStyle, NativeSpreadsheetDataValidationOperator, - NativeSpreadsheetDataValidationType, NativeSpreadsheetDifferentialFormat, + NativeSpreadsheetDataValidationType, NativeSpreadsheetDelimitedFormat, + NativeSpreadsheetDelimitedImport, NativeSpreadsheetDifferentialFormat, NativeSpreadsheetDynamicFilter, NativeSpreadsheetFill, NativeSpreadsheetFilterColumn, - NativeSpreadsheetFilterCriteria, NativeSpreadsheetNamedRange, NativeSpreadsheetNamedRangeScope, - NativeSpreadsheetReadingOrder, NativeSpreadsheetSort, NativeSpreadsheetSortDirection, - NativeSpreadsheetSortKey, NativeSpreadsheetTable, NativeSpreadsheetTableColumn, - NativeSpreadsheetTableStyle, NativeSpreadsheetVerticalAlignment, SpreadsheetCellValue, - MAX_NATIVE_OFFICE_FIND_BYTES, MAX_NATIVE_OFFICE_REPLACEMENT_BYTES, - MAX_NATIVE_OFFICE_TEXT_MATCHES, MAX_NATIVE_OFFICE_TEXT_REPLACEMENT_OUTPUT_BYTES, - MAX_NATIVE_OFFICE_TEXT_SCOPE_CELLS, + NativeSpreadsheetFilterCriteria, NativeSpreadsheetFrozenPane, NativeSpreadsheetImportResult, + NativeSpreadsheetNamedRange, NativeSpreadsheetNamedRangeScope, NativeSpreadsheetReadingOrder, + NativeSpreadsheetSort, NativeSpreadsheetSortDirection, NativeSpreadsheetSortKey, + NativeSpreadsheetTable, NativeSpreadsheetTableColumn, NativeSpreadsheetTableStyle, + NativeSpreadsheetVerticalAlignment, SpreadsheetCellValue, MAX_NATIVE_OFFICE_FIND_BYTES, + MAX_NATIVE_OFFICE_REPLACEMENT_BYTES, MAX_NATIVE_OFFICE_TEXT_MATCHES, + MAX_NATIVE_OFFICE_TEXT_REPLACEMENT_OUTPUT_BYTES, MAX_NATIVE_OFFICE_TEXT_SCOPE_CELLS, + MAX_NATIVE_SPREADSHEET_IMPORT_BYTES, MAX_NATIVE_SPREADSHEET_IMPORT_CELLS, }; pub use install::{install_office_cli, repair_office_cli, uninstall_managed_office_cli}; pub use issues::{ @@ -89,6 +91,23 @@ pub use semantic::{ NativeOfficeAnnotatedView, NativeOfficeDocument, OfficeNodeType, OutlineEntry, TextBlock, TextView, DEFAULT_NATIVE_OFFICE_ANNOTATED_LIMIT, MAX_NATIVE_OFFICE_ANNOTATED_LIMIT, }; +pub use spreadsheet_formula::{ + parse_spreadsheet_formula, SpreadsheetFormula, SpreadsheetFormulaBinaryOperator, + SpreadsheetFormulaCalculatedCell, SpreadsheetFormulaCalculation, SpreadsheetFormulaCell, + SpreadsheetFormulaDependencyGraph, SpreadsheetFormulaDependencyNode, + SpreadsheetFormulaErrorLiteral, SpreadsheetFormulaExpression, SpreadsheetFormulaExpressionKind, + SpreadsheetFormulaFunctionDefinition, SpreadsheetFormulaFunctionRegistry, + SpreadsheetFormulaFunctionReturnKind, SpreadsheetFormulaFunctionVolatility, + SpreadsheetFormulaLiteral, SpreadsheetFormulaPostfixOperator, SpreadsheetFormulaQualifier, + SpreadsheetFormulaReference, SpreadsheetFormulaReferenceKind, SpreadsheetFormulaSpan, + SpreadsheetFormulaUnaryOperator, SpreadsheetFormulaUnresolvedReference, + SpreadsheetFormulaUnresolvedReferenceKind, SpreadsheetFormulaValue, + MAX_SPREADSHEET_FORMULA_CALCULATION_TEXT_BYTES, MAX_SPREADSHEET_FORMULA_CELLS, + MAX_SPREADSHEET_FORMULA_CHARACTERS, MAX_SPREADSHEET_FORMULA_DEPENDENCIES, + MAX_SPREADSHEET_FORMULA_DEPTH, MAX_SPREADSHEET_FORMULA_NODES, + MAX_SPREADSHEET_FORMULA_REFERENCE_AREAS, MAX_SPREADSHEET_FORMULA_REFERENCE_VISITS, + MAX_SPREADSHEET_FORMULA_SPILL_CELLS, MAX_SPREADSHEET_FORMULA_TEXT_BYTES, +}; pub use template_merge::{ NativeOfficeTemplateMergeResult, MAX_TEMPLATE_DATA_DEPTH, MAX_TEMPLATE_DATA_ENTRIES, MAX_TEMPLATE_DATA_FLATTENED_BYTES, MAX_TEMPLATE_DATA_KEY_BYTES, @@ -305,6 +324,14 @@ mod spreadsheet_edit_tests; #[cfg(test)] mod spreadsheet_filter_tests; +#[cfg(test)] +mod spreadsheet_formula_calculation_tests; + +#[cfg(test)] +mod spreadsheet_formula_graph_tests; + +#[cfg(test)] +mod spreadsheet_import_tests; #[cfg(test)] mod spreadsheet_sort_tests; diff --git a/crates/office/src/replay.rs b/crates/office/src/replay.rs index de7783fc..4cff977e 100644 --- a/crates/office/src/replay.rs +++ b/crates/office/src/replay.rs @@ -6,13 +6,14 @@ use serde::{Deserialize, Serialize}; use crate::discovery::office_error; use crate::editor::{ NativeOfficeEditor, NativeOfficeMutation, NativeSpreadsheetAutoFilter, - NativeSpreadsheetConditionalFormat, NativeSpreadsheetDataValidation, - NativeSpreadsheetDataValidationErrorStyle, NativeSpreadsheetDataValidationOperator, - NativeSpreadsheetDataValidationType, NativeSpreadsheetNamedRange, - NativeSpreadsheetNamedRangeScope, NativeSpreadsheetSort, NativeSpreadsheetTable, - SpreadsheetCellValue, + NativeSpreadsheetCellFormat, NativeSpreadsheetConditionalFormat, + NativeSpreadsheetDataValidation, NativeSpreadsheetDataValidationErrorStyle, + NativeSpreadsheetDataValidationOperator, NativeSpreadsheetDataValidationType, + NativeSpreadsheetFrozenPane, NativeSpreadsheetNamedRange, NativeSpreadsheetNamedRangeScope, + NativeSpreadsheetSort, NativeSpreadsheetTable, SpreadsheetCellValue, }; use crate::semantic::{DocumentNode, NativeOfficeDocument, OfficeNodeType}; +use crate::spreadsheet_reference::{CellRange, CellReference}; use crate::{DocumentKind, NativeOfficePackage}; pub const NATIVE_OFFICE_REPLAY_FORMAT: &str = "a3s.office.native-replay"; @@ -327,6 +328,7 @@ fn emit_spreadsheet(root: &DocumentNode) -> UseResult> )); } let mut mutations = Vec::new(); + let mut requires_recalculation = false; for (sheet_offset, sheet) in sheets.into_iter().enumerate() { let name = sheet .path @@ -351,15 +353,37 @@ fn emit_spreadsheet(root: &DocumentNode) -> UseResult> let mut merged_ranges = Vec::new(); let mut auto_filter = None; let mut sort_state = None; + let mut frozen_pane = None; + let mut cell_formats = Vec::new(); let mut conditional_formats = Vec::new(); let mut tables = Vec::new(); let mut validations = Vec::new(); + let spill_owners = spreadsheet_spill_owners(sheet)?; for child in &sheet.children { match child.node_type { OfficeNodeType::Row => { require_plain_node(child, OfficeNodeType::Row)?; for cell in &child.children { - require_spreadsheet_cell(cell)?; + let reference = spreadsheet_cell_reference(cell)?; + if spill_owners + .get(&reference) + .is_some_and(|anchor| *anchor != reference) + { + if cell.format.contains_key("formula") { + return Err(dump_unsupported( + &cell.path, + "Legacy multi-cell array formula storage is not replayable.", + )); + } + continue; + } + if let Some(format) = spreadsheet_cell_format(cell)? { + cell_formats.push((cell.path.clone(), format)); + } + requires_recalculation |= cell + .format + .get("formulaCached") + .is_some_and(|cached| cached == "true"); mutations.push(NativeOfficeMutation::SetCellValue { path: cell.path.clone(), value: spreadsheet_value(cell)?, @@ -460,6 +484,30 @@ fn emit_spreadsheet(root: &DocumentNode) -> UseResult> })?; sort_state = Some((format!("{}/{}", sheet.path, reference), sort)); } + OfficeNodeType::FrozenPane => { + if frozen_pane.is_some() { + return Err(dump_unsupported( + &child.path, + "Spreadsheet replay found multiple frozen panes.", + )); + } + if child.format.get("nativeMutable").map(String::as_str) != Some("true") { + return Err(dump_unsupported( + &child.path, + "Spreadsheet frozen pane contains unsupported view state.", + )); + } + frozen_pane = Some( + NativeSpreadsheetFrozenPane::from_semantic_node(child).map_err( + |error| { + dump_unsupported( + &child.path, + format!("Spreadsheet frozen pane is not replayable: {error}"), + ) + }, + )?, + ); + } _ => { return Err(dump_unsupported( &child.path, @@ -473,6 +521,11 @@ fn emit_spreadsheet(root: &DocumentNode) -> UseResult> .into_iter() .map(|path| NativeOfficeMutation::MergeCells { path }), ); + mutations.extend( + cell_formats + .into_iter() + .map(|(path, format)| NativeOfficeMutation::SetCellFormat { path, format }), + ); mutations.extend(conditional_formats.into_iter().map(|conditional_format| { NativeOfficeMutation::AddConditionalFormat { sheet: sheet.path.clone(), @@ -491,6 +544,12 @@ fn emit_spreadsheet(root: &DocumentNode) -> UseResult> filter, }); } + if let Some(pane) = frozen_pane { + mutations.push(NativeOfficeMutation::SetSpreadsheetFrozenPane { + sheet: sheet.path.clone(), + pane, + }); + } mutations.extend(tables.into_iter().map(|table| { NativeOfficeMutation::AddSpreadsheetTable { sheet: sheet.path.clone(), @@ -540,9 +599,110 @@ fn emit_spreadsheet(root: &DocumentNode) -> UseResult> "Spreadsheet root contains an unsupported semantic node.", )); } + if requires_recalculation { + mutations.push(NativeOfficeMutation::RecalculateSpreadsheetFormulas); + } Ok(mutations) } +fn spreadsheet_spill_owners( + sheet: &DocumentNode, +) -> UseResult> { + let mut owners = std::collections::BTreeMap::new(); + for cell in sheet + .children + .iter() + .filter(|node| node.node_type == OfficeNodeType::Row) + .flat_map(|row| &row.children) + .filter(|node| node.node_type == OfficeNodeType::Cell) + { + let formula_type = cell.format.get("formulaType").map(String::as_str); + let formula_reference = cell.format.get("formulaRef"); + match (formula_type, formula_reference) { + (None, None) => continue, + (Some(kind), None) => { + return Err(dump_unsupported( + &cell.path, + format!( + "Spreadsheet formula storage type '{kind}' is not canonical replay input." + ), + )); + } + (Some(kind), Some(_)) if !kind.eq_ignore_ascii_case("array") => { + return Err(dump_unsupported( + &cell.path, + format!( + "Spreadsheet formula storage type '{kind}' with a spill range is not replayable." + ), + )); + } + (None, Some(_)) => { + return Err(dump_unsupported( + &cell.path, + "Spreadsheet formula spill range has no array storage type.", + )); + } + (Some(_), Some(_)) => {} + } + if cell.format.get("formulaCached").map(String::as_str) != Some("true") { + return Err(dump_unsupported( + &cell.path, + "Spreadsheet array formula has no cached native result.", + )); + } + let reference = formula_reference.ok_or_else(|| { + dump_unsupported(&cell.path, "Spreadsheet array formula has no spill range.") + })?; + let anchor = spreadsheet_cell_reference(cell)?; + let range = CellRange::parse(reference).map_err(|error| { + dump_unsupported( + &cell.path, + format!("Spreadsheet formula spill range '{reference}' is invalid: {error}"), + ) + })?; + if !range.contains(anchor) { + return Err(dump_unsupported( + &cell.path, + "Spreadsheet formula spill range does not contain its anchor.", + )); + } + let cells = range.cell_count()?; + if cells > crate::MAX_SPREADSHEET_FORMULA_SPILL_CELLS + || owners.len().saturating_add(cells) > crate::MAX_SPREADSHEET_FORMULA_SPILL_CELLS + { + return Err(dump_unsupported( + &cell.path, + "Spreadsheet formula spills exceed the native replay cell limit.", + ) + .with_detail("cells", owners.len().saturating_add(cells))); + } + for row in range.start.row..=range.end.row { + for column in range.start.column..=range.end.column { + let reference = CellReference { column, row }; + if owners.insert(reference, anchor).is_some() { + return Err(dump_unsupported( + &cell.path, + "Spreadsheet formula spill ranges overlap.", + )); + } + } + } + } + Ok(owners) +} + +fn spreadsheet_cell_reference(cell: &DocumentNode) -> UseResult { + cell.path + .rsplit_once('/') + .and_then(|(_, reference)| CellReference::parse(reference).ok()) + .ok_or_else(|| { + dump_unsupported( + &cell.path, + "Spreadsheet cell has an invalid semantic coordinate.", + ) + }) +} + fn spreadsheet_named_range(node: &DocumentNode) -> UseResult { if node.node_type != OfficeNodeType::NamedRange || node.tag != "namedrange" @@ -723,45 +883,105 @@ fn parse_bool_format(node: &DocumentNode, key: &str, default: bool) -> UseResult } } -fn require_spreadsheet_cell(cell: &DocumentNode) -> UseResult<()> { +fn spreadsheet_cell_format(cell: &DocumentNode) -> UseResult> { if cell.node_type != OfficeNodeType::Cell || cell.style.is_some() || !cell.children.is_empty() { return Err(dump_unsupported( &cell.path, "Spreadsheet replay requires a leaf cell without semantic child nodes.", )); } - let allowed = [ + let plain = [ "column", "row", "valueType", + "valuePresent", "empty", "formula", + "formulaCached", + "formulaRef", + "formulaType", "merge", "mergeAnchor", "dataValidation", "validationType", ]; + let styled = [ + "styleIndex", + "baseStyleId", + "fontId", + "font", + "size", + "fillId", + "borderId", + "numberFormatId", + "numberFormat", + ]; if let Some(key) = cell .format .keys() - .find(|key| !allowed.contains(&key.as_str())) + .find(|key| !plain.contains(&key.as_str()) && !styled.contains(&key.as_str())) { return Err(dump_unsupported( &cell.path, format!("Spreadsheet cell property '{key}' is not replayable yet."), )); } - Ok(()) -} - -fn spreadsheet_value(cell: &DocumentNode) -> UseResult { - if let Some(expression) = cell.format.get("formula") { - if !cell.text.is_empty() { + if !cell.format.contains_key("styleIndex") { + if let Some(key) = styled.iter().find(|key| cell.format.contains_key(**key)) { return Err(dump_unsupported( &cell.path, - "Spreadsheet formulas with cached results are not exactly replayable yet.", + format!("Spreadsheet cell has style property '{key}' without a styleIndex."), )); } + return Ok(None); + } + for (key, expected) in [ + ("baseStyleId", "0"), + ("fontId", "0"), + ("font", "Aptos"), + ("size", "11pt"), + ("fillId", "0"), + ("borderId", "0"), + ] { + if cell.format.get(key).map(String::as_str) != Some(expected) { + return Err(dump_unsupported( + &cell.path, + format!( + "Spreadsheet replay supports only the native default date style; '{key}' is not '{expected}'." + ), + )); + } + } + let number_format = cell.format.get("numberFormat").ok_or_else(|| { + dump_unsupported( + &cell.path, + "Spreadsheet styled cell has no numberFormat value.", + ) + })?; + let valid_number_format_id = cell + .format + .get("numberFormatId") + .and_then(|value| value.parse::().ok()) + .is_some_and(|value| value >= 164); + let valid_style_index = cell + .format + .get("styleIndex") + .and_then(|value| value.parse::().ok()) + .is_some(); + if number_format != "yyyy-mm-dd" || !valid_number_format_id || !valid_style_index { + return Err(dump_unsupported( + &cell.path, + "Spreadsheet replay supports only the canonical native yyyy-mm-dd import style.", + )); + } + Ok(Some(NativeSpreadsheetCellFormat { + number_format: Some(number_format.clone()), + ..NativeSpreadsheetCellFormat::default() + })) +} + +fn spreadsheet_value(cell: &DocumentNode) -> UseResult { + if let Some(expression) = cell.format.get("formula") { return Ok(SpreadsheetCellValue::Formula { expression: expression.clone(), }); diff --git a/crates/office/src/semantic/mod.rs b/crates/office/src/semantic/mod.rs index f1e81770..cea28f1a 100644 --- a/crates/office/src/semantic/mod.rs +++ b/crates/office/src/semantic/mod.rs @@ -49,6 +49,7 @@ pub enum OfficeNodeType { FilterValue, SortState, SortKey, + FrozenPane, Presentation, Slide, Shape, @@ -90,6 +91,7 @@ impl OfficeNodeType { Self::FilterValue => "FilterValue", Self::SortState => "SortState", Self::SortKey => "SortKey", + Self::FrozenPane => "FrozenPane", Self::Presentation => "Presentation", Self::Slide => "Slide", Self::Shape => "Shape", diff --git a/crates/office/src/semantic/selector.rs b/crates/office/src/semantic/selector.rs index c10d2b5c..72a46db7 100644 --- a/crates/office/src/semantic/selector.rs +++ b/crates/office/src/semantic/selector.rs @@ -264,6 +264,7 @@ fn kind_matches(kind: &str, node: &DocumentNode) -> bool { node.node_type == OfficeNodeType::NamedRangeCollection } "sheet" | "worksheet" => node.node_type == OfficeNodeType::Worksheet, + "freeze" | "frozenpane" | "frozen-pane" => node.node_type == OfficeNodeType::FrozenPane, "slide" => node.node_type == OfficeNodeType::Slide, "shape" | "textbox" => matches!( node.node_type, diff --git a/crates/office/src/semantic/spreadsheet.rs b/crates/office/src/semantic/spreadsheet.rs index 030dfb82..a93bc49b 100644 --- a/crates/office/src/semantic/spreadsheet.rs +++ b/crates/office/src/semantic/spreadsheet.rs @@ -20,6 +20,7 @@ mod named_range; mod sort_state; mod style; mod table; +mod view; use style::{read_differential_formats, read_styles}; @@ -437,6 +438,9 @@ fn read_worksheet( if let Some(sort) = sort_state::read(&worksheet, part_name, &sheet_path)? { sheet_node.children.push(sort); } + if let Some(freeze) = view::read(&worksheet, &sheet_path) { + sheet_node.children.push(freeze); + } let tables = table::read(package, opc, &worksheet, part_name, &sheet_path)?; if !tables.is_empty() { sheet_node @@ -1092,8 +1096,16 @@ fn read_cell( .insert("valueType".into(), value_type_name(value_type).into()); node.format .insert("empty".into(), node.text.is_empty().to_string()); + node.format.insert( + "valuePresent".into(), + (cell.child("v").is_some() || cell.child("is").is_some()).to_string(), + ); if let Some(formula) = cell.child("f") { node.format.insert("formula".into(), direct_text(formula)); + node.format.insert( + "formulaCached".into(), + cell.child("v").is_some().to_string(), + ); if let Some(formula_type) = formula.attribute("t") { node.format .insert("formulaType".into(), formula_type.into()); diff --git a/crates/office/src/semantic/spreadsheet/view.rs b/crates/office/src/semantic/spreadsheet/view.rs new file mode 100644 index 00000000..cec21ade --- /dev/null +++ b/crates/office/src/semantic/spreadsheet/view.rs @@ -0,0 +1,90 @@ +use super::{DocumentNode, OfficeNodeType}; +use crate::xml_tree::{XmlElement, XmlNode}; +use crate::NativeSpreadsheetFrozenPane; + +pub(super) fn read(worksheet: &XmlElement, sheet_path: &str) -> Option { + let views = worksheet + .children_named("sheetViews") + .filter(|views| views.namespace == worksheet.namespace) + .collect::>(); + let collection = *views.first()?; + let sheet_views = collection + .children_named("sheetView") + .filter(|view| view.namespace == worksheet.namespace) + .collect::>(); + let view = sheet_views + .iter() + .copied() + .find(|view| view.attribute("workbookViewId") == Some("0")) + .or_else(|| sheet_views.first().copied())?; + let panes = view + .children_named("pane") + .filter(|pane| pane.namespace == worksheet.namespace) + .collect::>(); + let pane = *panes.first()?; + let frozen_rows = integer_split(pane.attribute("ySplit")); + let frozen_columns = integer_split(pane.attribute("xSplit")); + let top_left = pane.attribute("topLeftCell").unwrap_or_default(); + let candidate = NativeSpreadsheetFrozenPane::new( + frozen_rows.unwrap_or_default(), + frozen_columns.unwrap_or_default(), + top_left, + ); + let normalized = candidate.normalized().ok(); + let expected_active = normalized + .as_ref() + .map(NativeSpreadsheetFrozenPane::active_pane); + let known_attributes = pane.attributes.iter().all(|attribute| { + attribute.namespace.is_none() + && matches!( + attribute.local_name.as_str(), + "xSplit" | "ySplit" | "topLeftCell" | "activePane" | "state" + ) + }); + let empty = pane.children.iter().all(|child| match child { + XmlNode::Text(text) => text.trim().is_empty(), + XmlNode::Element(_) => false, + }); + let mutable = views.len() == 1 + && view.attribute("workbookViewId") == Some("0") + && panes.len() == 1 + && frozen_rows.is_some() + && frozen_columns.is_some() + && normalized.is_some() + && pane.attribute("state") == Some("frozen") + && pane.attribute("activePane") == expected_active + && known_attributes + && empty; + + let mut node = DocumentNode::new( + format!("{sheet_path}/freeze"), + "freeze", + OfficeNodeType::FrozenPane, + ); + node.text = format!( + "{} frozen row(s), {} frozen column(s)", + candidate.frozen_rows, candidate.frozen_columns + ); + node.format + .insert("frozenRows".into(), candidate.frozen_rows.to_string()); + node.format + .insert("frozenColumns".into(), candidate.frozen_columns.to_string()); + node.format + .insert("topLeftCell".into(), candidate.top_left_cell); + if let Some(active) = pane.attribute("activePane") { + node.format.insert("activePane".into(), active.into()); + } + if let Some(state) = pane.attribute("state") { + node.format.insert("state".into(), state.into()); + } + node.format + .insert("nativeMutable".into(), mutable.to_string()); + Some(node) +} + +fn integer_split(value: Option<&str>) -> Option { + match value { + None => Some(0), + Some(value) => value.parse::().ok(), + } +} diff --git a/crates/office/src/semantic_tests.rs b/crates/office/src/semantic_tests.rs index 6ed0ce46..8457bdb0 100644 --- a/crates/office/src/semantic_tests.rs +++ b/crates/office/src/semantic_tests.rs @@ -412,6 +412,35 @@ async fn native_editor_writes_typed_spreadsheet_values_and_marks_formulas_for_re .contains("calcChain")); } +#[tokio::test] +async fn native_editor_rejects_syntactically_invalid_formulas_atomically() { + let temp = tempfile::tempdir().unwrap(); + let path = temp.path().join("invalid-formulas.xlsx"); + let mut editor = NativeOfficeEditor::create(&path).await.unwrap(); + let original = editor.package().content_sha256(); + + for expression in ["1+", "SUM(A1", "A1::B2", "\"unterminated"] { + let error = editor + .set_cell_value( + "/Sheet1/A1", + SpreadsheetCellValue::Formula { + expression: expression.into(), + }, + ) + .unwrap_err(); + assert_eq!( + error.code, "use.office.spreadsheet_formula_invalid", + "{expression}" + ); + assert!( + error.details.contains_key("characterOffset"), + "{expression}" + ); + assert!(error.details.contains_key("byteOffset"), "{expression}"); + assert_eq!(editor.package().content_sha256(), original, "{expression}"); + } +} + #[tokio::test] async fn native_editor_adds_a_worksheet_and_populates_it() { let temp = tempfile::tempdir().unwrap(); diff --git a/crates/office/src/spreadsheet_formula.rs b/crates/office/src/spreadsheet_formula.rs index f2cfaa81..edd03475 100644 --- a/crates/office/src/spreadsheet_formula.rs +++ b/crates/office/src/spreadsheet_formula.rs @@ -3,6 +3,139 @@ use a3s_use_core::{UseError, UseResult}; use crate::discovery::office_error; use crate::spreadsheet_reference::{column_name, MAX_COLUMNS, MAX_ROWS}; +mod ast; +mod evaluate; +mod graph; +mod lexer; +mod parser; +mod registry; +mod structured_reference; +mod value; + +pub use ast::{ + SpreadsheetFormula, SpreadsheetFormulaBinaryOperator, SpreadsheetFormulaErrorLiteral, + SpreadsheetFormulaExpression, SpreadsheetFormulaExpressionKind, SpreadsheetFormulaLiteral, + SpreadsheetFormulaPostfixOperator, SpreadsheetFormulaQualifier, SpreadsheetFormulaReference, + SpreadsheetFormulaReferenceKind, SpreadsheetFormulaSpan, SpreadsheetFormulaUnaryOperator, + MAX_SPREADSHEET_FORMULA_CHARACTERS, MAX_SPREADSHEET_FORMULA_DEPTH, + MAX_SPREADSHEET_FORMULA_NODES, MAX_SPREADSHEET_FORMULA_REFERENCE_AREAS, +}; +pub use evaluate::{ + MAX_SPREADSHEET_FORMULA_CALCULATION_TEXT_BYTES, MAX_SPREADSHEET_FORMULA_SPILL_CELLS, + MAX_SPREADSHEET_FORMULA_TEXT_BYTES, +}; +pub use graph::{ + SpreadsheetFormulaCell, SpreadsheetFormulaDependencyGraph, SpreadsheetFormulaDependencyNode, + SpreadsheetFormulaUnresolvedReference, SpreadsheetFormulaUnresolvedReferenceKind, + MAX_SPREADSHEET_FORMULA_CELLS, MAX_SPREADSHEET_FORMULA_DEPENDENCIES, + MAX_SPREADSHEET_FORMULA_REFERENCE_VISITS, +}; +pub use registry::{ + SpreadsheetFormulaFunctionDefinition, SpreadsheetFormulaFunctionRegistry, + SpreadsheetFormulaFunctionReturnKind, SpreadsheetFormulaFunctionVolatility, +}; +pub(crate) use structured_reference::{ + LocalStructuredReferenceContext, StructuredReferenceRewritePlan, + StructuredReferenceRewriteResult, +}; +pub use value::{ + SpreadsheetFormulaCalculatedCell, SpreadsheetFormulaCalculation, SpreadsheetFormulaValue, +}; + +#[derive(Debug)] +struct FormulaParseFailure { + byte_offset: usize, + reason: String, +} + +impl FormulaParseFailure { + fn new(byte_offset: usize, reason: impl Into) -> Self { + Self { + byte_offset, + reason: reason.into(), + } + } +} + +/// Parses one Spreadsheet formula into a bounded, source-spanned typed AST. +/// +/// Callers may provide a formula-bar leading `=`. Spans and parse-error +/// positions address the normalized formula body after that optional marker. +pub fn parse_spreadsheet_formula(formula: &str) -> UseResult { + let normalized = formula.strip_prefix('=').unwrap_or(formula); + parse_normalized_formula(normalized) +} + +pub(crate) fn validate_and_normalize_formula(formula: &str) -> UseResult<&str> { + let normalized = formula.strip_prefix('=').unwrap_or(formula); + parse_normalized_formula(normalized)?; + Ok(normalized) +} + +fn parse_normalized_formula(formula: &str) -> UseResult { + validate_formula_bounds(formula)?; + let tokens = lexer::lex(formula).map_err(|failure| parse_error(formula, failure))?; + parser::parse(formula, tokens).map_err(|failure| parse_error(formula, failure)) +} + +fn validate_formula_bounds(formula: &str) -> UseResult<()> { + let characters = formula.chars().count(); + if formula.is_empty() || characters > MAX_SPREADSHEET_FORMULA_CHARACTERS { + return Err(office_error( + "use.office.spreadsheet_formula_invalid", + format!( + "Spreadsheet formulas must contain 1-{MAX_SPREADSHEET_FORMULA_CHARACTERS} characters." + ), + ) + .with_detail("characterOffset", characters) + .with_detail("byteOffset", formula.len()) + .with_detail("reason", "Formula length is outside supported limits.")); + } + if let Some((byte_offset, _)) = formula + .char_indices() + .find(|(_, character)| character.is_control()) + { + let character_offset = formula[..byte_offset].chars().count(); + return Err(office_error( + "use.office.spreadsheet_formula_invalid", + format!( + "Spreadsheet formula is invalid at character {}: control characters are not supported.", + character_offset + 1 + ), + ) + .with_detail("characterOffset", character_offset) + .with_detail("byteOffset", byte_offset) + .with_detail( + "reason", + "Formula contains an unsupported control character.", + )); + } + Ok(()) +} + +fn parse_error(formula: &str, failure: FormulaParseFailure) -> UseError { + let byte_offset = nearest_character_boundary(formula, failure.byte_offset.min(formula.len())); + let character_offset = formula[..byte_offset].chars().count(); + office_error( + "use.office.spreadsheet_formula_invalid", + format!( + "Spreadsheet formula is invalid at character {}: {}", + character_offset + 1, + failure.reason + ), + ) + .with_detail("characterOffset", character_offset) + .with_detail("byteOffset", byte_offset) + .with_detail("reason", failure.reason) +} + +fn nearest_character_boundary(value: &str, mut offset: usize) -> usize { + while offset > 0 && !value.is_char_boundary(offset) { + offset -= 1; + } + offset +} + #[derive(Debug, Clone, Copy, PartialEq, Eq)] pub(crate) enum ReferenceAxis { Row, @@ -367,7 +500,9 @@ fn rewrite_reference( ParsedReference::Rows { start, end } => { rewrite_axis_reference(*start, *end, ReferenceAxis::Row, axis, edit) } - ParsedReference::Cells { .. } => unreachable!(), + ParsedReference::Cells { .. } => Err(formula_error( + "Spreadsheet formula reference classification is inconsistent.", + )), }; }; let Some(end) = end else { @@ -599,6 +734,68 @@ fn formula_error(message: impl Into) -> UseError { mod tests { use super::*; + #[test] + fn public_formula_parser_is_bounded_utf8_safe_and_send_sync() { + fn assert_send_sync() {} + assert_send_sync::(); + + let formula = parse_spreadsheet_formula("=SUM('销售 数据'!A1:B2,Table1[Amount])").unwrap(); + assert!(matches!( + formula.root.kind, + SpreadsheetFormulaExpressionKind::FunctionCall { .. } + )); + assert!(parse_spreadsheet_formula("==1").is_err()); + + let error = parse_spreadsheet_formula("\"销售\"+").unwrap_err(); + assert_eq!(error.code, "use.office.spreadsheet_formula_invalid"); + assert_eq!(error.details["characterOffset"], 5); + assert_eq!(error.details["byteOffset"], 9); + + let too_long = "1".repeat(MAX_SPREADSHEET_FORMULA_CHARACTERS + 1); + assert_eq!( + parse_spreadsheet_formula(&too_long).unwrap_err().code, + "use.office.spreadsheet_formula_invalid" + ); + let too_deep = format!( + "{}1{}", + "(".repeat(MAX_SPREADSHEET_FORMULA_DEPTH + 1), + ")".repeat(MAX_SPREADSHEET_FORMULA_DEPTH + 1) + ); + assert_eq!( + parse_spreadsheet_formula(&too_deep).unwrap_err().code, + "use.office.spreadsheet_formula_invalid" + ); + } + + #[test] + fn formula_parser_rejects_incomplete_and_non_excel_operator_syntax() { + for formula in [ + "", + " ", + "1+", + "+", + "SUM(A1", + "SUM(A1;B1)", + "A1::B2", + "A:1", + "1 2", + "A1&&B1", + "A1!=B1", + "\"unterminated", + "'unterminated!A1", + "Table1[[Column]", + "$A+1", + ] { + let error = parse_spreadsheet_formula(formula).unwrap_err(); + assert_eq!( + error.code, "use.office.spreadsheet_formula_invalid", + "{formula}" + ); + assert!(error.details.contains_key("characterOffset"), "{formula}"); + assert!(error.details.contains_key("byteOffset"), "{formula}"); + } + } + #[test] fn structural_rewrite_respects_sheets_strings_ranges_and_absolute_markers() { let formula = r#"SUM(A1,$B$2,A3:A5,'Data Set'!C4,"A1",Other!D6)"#; diff --git a/crates/office/src/spreadsheet_formula/ast.rs b/crates/office/src/spreadsheet_formula/ast.rs new file mode 100644 index 00000000..3767dc2c --- /dev/null +++ b/crates/office/src/spreadsheet_formula/ast.rs @@ -0,0 +1,245 @@ +/// Maximum number of Unicode scalar values accepted in one cell formula. +pub const MAX_SPREADSHEET_FORMULA_CHARACTERS: usize = 8_192; + +/// Maximum nesting depth accepted by the native formula parser. +pub const MAX_SPREADSHEET_FORMULA_DEPTH: usize = 128; + +/// Maximum number of AST nodes accepted by the native formula parser. +pub const MAX_SPREADSHEET_FORMULA_NODES: usize = 8_192; + +/// Maximum disjoint reference areas retained while graphing or calculating +/// one Spreadsheet formula value. +pub const MAX_SPREADSHEET_FORMULA_REFERENCE_AREAS: usize = 100_000; + +/// A parsed Spreadsheet formula body. +/// +/// The source may include one formula-bar leading `=`. The parser removes that +/// marker before assigning spans, so all spans address the normalized formula +/// body that is stored in SpreadsheetML. +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct SpreadsheetFormula { + pub root: SpreadsheetFormulaExpression, +} + +/// A half-open UTF-8 byte range in the normalized formula body. +#[derive(Debug, Default, Clone, Copy, PartialEq, Eq)] +pub struct SpreadsheetFormulaSpan { + pub start: usize, + pub end: usize, +} + +impl SpreadsheetFormulaSpan { + pub(crate) const fn new(start: usize, end: usize) -> Self { + Self { start, end } + } + + pub(crate) const fn through(self, other: Self) -> Self { + Self { + start: self.start, + end: other.end, + } + } +} + +/// One source-spanned formula expression. +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct SpreadsheetFormulaExpression { + pub span: SpreadsheetFormulaSpan, + pub kind: SpreadsheetFormulaExpressionKind, +} + +/// Closed expression variants produced by the native Spreadsheet parser. +#[derive(Debug, Clone, PartialEq, Eq)] +pub enum SpreadsheetFormulaExpressionKind { + Literal(SpreadsheetFormulaLiteral), + Reference(SpreadsheetFormulaReference), + Name { + qualifier: Option, + name: String, + }, + StructuredReference { + qualifier: Option, + reference: String, + }, + Unary { + operator: SpreadsheetFormulaUnaryOperator, + operand: Box, + }, + Postfix { + operator: SpreadsheetFormulaPostfixOperator, + operand: Box, + }, + Binary { + operator: SpreadsheetFormulaBinaryOperator, + left: Box, + right: Box, + }, + FunctionCall { + qualifier: Option, + name: String, + arguments: Vec>, + }, + Parenthesized(Box), + Array { + rows: Vec>, + }, +} + +/// Scalar literal values represented directly in formula source. +#[derive(Debug, Clone, PartialEq, Eq)] +pub enum SpreadsheetFormulaLiteral { + Number(String), + Text(String), + Boolean(bool), + Error(SpreadsheetFormulaErrorLiteral), +} + +/// Error literals supported by current Spreadsheet applications. +#[derive(Debug, Clone, Copy, PartialEq, Eq, serde::Serialize, serde::Deserialize)] +#[serde(rename_all = "kebab-case")] +pub enum SpreadsheetFormulaErrorLiteral { + Null, + DivisionByZero, + Value, + Reference, + Name, + Number, + NotAvailable, + GettingData, + Spill, + Calculation, + Field, + Blocked, + Unknown, + Busy, + Connect, + Python, +} + +impl SpreadsheetFormulaErrorLiteral { + pub fn parse(value: &str) -> Option { + [ + Self::GettingData, + Self::DivisionByZero, + Self::NotAvailable, + Self::Calculation, + Self::Reference, + Self::Blocked, + Self::Unknown, + Self::Connect, + Self::Python, + Self::Value, + Self::Field, + Self::Spill, + Self::Number, + Self::Name, + Self::Busy, + Self::Null, + ] + .into_iter() + .find(|literal| value.eq_ignore_ascii_case(literal.as_str())) + } + + pub fn as_str(self) -> &'static str { + match self { + Self::Null => "#NULL!", + Self::DivisionByZero => "#DIV/0!", + Self::Value => "#VALUE!", + Self::Reference => "#REF!", + Self::Name => "#NAME?", + Self::Number => "#NUM!", + Self::NotAvailable => "#N/A", + Self::GettingData => "#GETTING_DATA", + Self::Spill => "#SPILL!", + Self::Calculation => "#CALC!", + Self::Field => "#FIELD!", + Self::Blocked => "#BLOCKED!", + Self::Unknown => "#UNKNOWN!", + Self::Busy => "#BUSY!", + Self::Connect => "#CONNECT!", + Self::Python => "#PYTHON!", + } + } +} + +/// Workbook and worksheet qualification attached to a reference or name. +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct SpreadsheetFormulaQualifier { + /// External workbook prefix, including its square brackets when present. + pub workbook: Option, + /// First worksheet in the qualifier. + pub worksheet: String, + /// Last worksheet for a three-dimensional worksheet qualifier. + pub worksheet_end: Option, +} + +impl SpreadsheetFormulaQualifier { + pub fn is_external(&self) -> bool { + self.workbook.is_some() + } + + pub fn is_three_dimensional(&self) -> bool { + self.worksheet_end.is_some() + } +} + +/// One absolute/relative A1 reference endpoint. +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct SpreadsheetFormulaReference { + pub qualifier: Option, + pub kind: SpreadsheetFormulaReferenceKind, +} + +/// A cell, whole-column, or whole-row A1 reference endpoint. +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum SpreadsheetFormulaReferenceKind { + Cell { + column: u32, + row: u32, + absolute_column: bool, + absolute_row: bool, + }, + Column { + column: u32, + absolute: bool, + }, + Row { + row: u32, + absolute: bool, + }, +} + +/// Prefix operators in Spreadsheet formula order. +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum SpreadsheetFormulaUnaryOperator { + Positive, + Negative, + ImplicitIntersection, +} + +/// Postfix operators in Spreadsheet formula order. +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum SpreadsheetFormulaPostfixOperator { + Percent, + Spill, +} + +/// Binary arithmetic, comparison, concatenation, and reference operators. +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum SpreadsheetFormulaBinaryOperator { + Range, + Intersection, + Union, + Power, + Multiply, + Divide, + Add, + Subtract, + Concatenate, + Equal, + NotEqual, + LessThan, + LessThanOrEqual, + GreaterThan, + GreaterThanOrEqual, +} diff --git a/crates/office/src/spreadsheet_formula/evaluate.rs b/crates/office/src/spreadsheet_formula/evaluate.rs new file mode 100644 index 00000000..97f800f5 --- /dev/null +++ b/crates/office/src/spreadsheet_formula/evaluate.rs @@ -0,0 +1,415 @@ +mod context; +mod function; +mod operators; +mod reference; +mod spill; + +use context::{ + build_context, public_cell_key, EvalValue, EvaluationContext, FormulaArray, FormulaCellKey, + FormulaReferenceArea, ScalarValue, +}; +use operators::{ + accumulate_formula_text_bytes, broadcast_dimension, calculation_error, + calculation_text_limit_error, checked_array_cells, ensure_formula_text_limit, + finite_or_number_error, format_number, into_array, invalid_array_shape, public_scalar, + scalar_binary, spill_limit_error, +}; + +use std::collections::BTreeMap; + +use a3s_use_core::UseResult; + +use crate::semantic::NativeOfficeDocument; +use crate::DocumentKind; + +use super::{ + SpreadsheetFormulaBinaryOperator, SpreadsheetFormulaCalculatedCell, + SpreadsheetFormulaCalculation, SpreadsheetFormulaCell, SpreadsheetFormulaErrorLiteral, + SpreadsheetFormulaExpression, SpreadsheetFormulaExpressionKind, + SpreadsheetFormulaFunctionRegistry, SpreadsheetFormulaLiteral, + SpreadsheetFormulaPostfixOperator, SpreadsheetFormulaUnaryOperator, +}; + +/// Maximum cells in one dynamic-array result and cumulative spill children in +/// one native calculation pass. +pub const MAX_SPREADSHEET_FORMULA_SPILL_CELLS: usize = 100_000; + +/// Maximum UTF-8 bytes produced by one native formula text result. +pub const MAX_SPREADSHEET_FORMULA_TEXT_BYTES: usize = 1024 * 1024; + +/// Maximum cumulative UTF-8 text-result bytes in one native calculation pass. +pub const MAX_SPREADSHEET_FORMULA_CALCULATION_TEXT_BYTES: usize = 8 * 1024 * 1024; + +impl NativeOfficeDocument { + /// Calculates supported formulas in memory without changing package bytes. + pub fn calculate_spreadsheet_formulas(&self) -> UseResult { + self.calculate_spreadsheet_formulas_with_registry( + &SpreadsheetFormulaFunctionRegistry::default(), + ) + } + + /// Calculates supported formulas using an explicit typed function registry. + pub fn calculate_spreadsheet_formulas_with_registry( + &self, + registry: &SpreadsheetFormulaFunctionRegistry, + ) -> UseResult { + calculate(self, registry) + } +} + +fn calculate( + document: &NativeOfficeDocument, + registry: &SpreadsheetFormulaFunctionRegistry, +) -> UseResult { + if document.kind() != DocumentKind::Spreadsheet { + return Err(calculation_error( + "use.office.spreadsheet_formula_calculation_type_unsupported", + "Native formula calculation is available only for Spreadsheet documents.", + )); + } + let graph = document.formula_dependency_graph()?; + if !graph.cycles.is_empty() { + return Err(calculation_error( + "use.office.spreadsheet_formula_cycle", + "Spreadsheet formula dependency graph contains a circular reference.", + ) + .with_detail( + "cycles", + serde_json::to_value( + graph + .cycles + .iter() + .map(|cycle| { + cycle + .iter() + .map(SpreadsheetFormulaCell::path) + .collect::>() + }) + .collect::>(), + ) + .unwrap_or(serde_json::Value::Null), + )); + } + let (mut context, records) = build_context(document, registry, &graph.nodes)?; + let record_indexes = records + .iter() + .enumerate() + .map(|(index, record)| (record.key, index)) + .collect::>(); + let mut cells = Vec::with_capacity(records.len()); + let mut spill_cell_count = 0_usize; + let mut text_result_bytes = 0_usize; + for cell in &graph.calculation_order { + let key = public_cell_key(cell, &context.sheet_names)?; + let record = record_indexes + .get(&key) + .and_then(|index| records.get(*index)) + .ok_or_else(|| { + calculation_error( + "use.office.spreadsheet_formula_calculation_invalid", + format!( + "Calculation order references missing formula cell '{}'.", + cell.path() + ), + ) + })?; + let evaluated = context + .evaluate_expression(&record.formula.root, key) + .map_err(|error| error.with_detail("cell", cell.path()))?; + let (value, spill_range, spill_cells, text_bytes) = + context.finalize_formula_result(key, evaluated, spill_cell_count, text_result_bytes)?; + spill_cell_count = spill_cell_count + .checked_add(spill_cells) + .ok_or_else(spill_limit_error)?; + text_result_bytes = text_result_bytes + .checked_add(text_bytes) + .ok_or_else(calculation_text_limit_error)?; + cells.push(SpreadsheetFormulaCalculatedCell { + cell: cell.clone(), + value, + spill_range, + }); + } + Ok(SpreadsheetFormulaCalculation { + formula_count: records.len(), + spill_cell_count, + calculation_order: graph.calculation_order, + cells, + }) +} + +impl EvaluationContext<'_> { + pub(super) fn evaluate_expression( + &mut self, + expression: &SpreadsheetFormulaExpression, + current: FormulaCellKey, + ) -> UseResult { + match &expression.kind { + SpreadsheetFormulaExpressionKind::Literal(literal) => { + Ok(EvalValue::Scalar(match literal { + SpreadsheetFormulaLiteral::Number(value) => value + .parse::() + .ok() + .filter(|value| value.is_finite()) + .map_or( + ScalarValue::Error(SpreadsheetFormulaErrorLiteral::Number), + ScalarValue::Number, + ), + SpreadsheetFormulaLiteral::Text(value) => ScalarValue::Text(value.clone()), + SpreadsheetFormulaLiteral::Boolean(value) => ScalarValue::Boolean(*value), + SpreadsheetFormulaLiteral::Error(error) => ScalarValue::Error(*error), + })) + } + SpreadsheetFormulaExpressionKind::Reference(reference) => { + self.evaluate_reference(reference, current) + } + SpreadsheetFormulaExpressionKind::Name { qualifier, name } => { + self.evaluate_name(qualifier.as_ref(), name, current) + } + SpreadsheetFormulaExpressionKind::StructuredReference { + qualifier, + reference, + } => self.evaluate_structured_reference(qualifier.as_ref(), reference, current), + SpreadsheetFormulaExpressionKind::Unary { operator, operand } => match operator { + SpreadsheetFormulaUnaryOperator::ImplicitIntersection => { + let value = self.evaluate_expression(operand, current)?; + self.implicit_intersection(value, current) + } + SpreadsheetFormulaUnaryOperator::Positive => { + let value = self.evaluate_expression(operand, current)?; + self.map_numeric(value, finite_or_number_error) + } + SpreadsheetFormulaUnaryOperator::Negative => { + let value = self.evaluate_expression(operand, current)?; + self.map_numeric(value, |number| finite_or_number_error(-number)) + } + }, + SpreadsheetFormulaExpressionKind::Postfix { operator, operand } => match operator { + SpreadsheetFormulaPostfixOperator::Percent => { + let value = self.evaluate_expression(operand, current)?; + self.map_numeric(value, |number| finite_or_number_error(number / 100.0)) + } + SpreadsheetFormulaPostfixOperator::Spill => { + let value = self.evaluate_expression(operand, current)?; + self.spill_reference(value) + } + }, + SpreadsheetFormulaExpressionKind::Binary { + operator, + left, + right, + } if matches!( + operator, + SpreadsheetFormulaBinaryOperator::Range + | SpreadsheetFormulaBinaryOperator::Intersection + | SpreadsheetFormulaBinaryOperator::Union + ) => + { + self.evaluate_reference_operator(*operator, left, right, current) + } + SpreadsheetFormulaExpressionKind::Binary { + operator, + left, + right, + } => { + let left = self.evaluate_expression(left, current)?; + let right = self.evaluate_expression(right, current)?; + self.apply_binary(*operator, left, right) + } + SpreadsheetFormulaExpressionKind::FunctionCall { + qualifier, + name, + arguments, + } => { + if qualifier.is_some() { + return Err(calculation_error( + "use.office.spreadsheet_formula_function_unsupported", + format!("Native calculation does not execute qualified function '{name}'."), + ) + .with_detail("function", name.clone())); + } + function::evaluate_function(self, name, arguments, current) + } + SpreadsheetFormulaExpressionKind::Parenthesized(inner) => { + self.evaluate_expression(inner, current) + } + SpreadsheetFormulaExpressionKind::Array { rows } => { + let mut values = Vec::with_capacity(rows.len()); + for row in rows { + let mut output = Vec::with_capacity(row.len()); + for value in row { + let evaluated = self.evaluate_expression(value, current)?; + output.push(self.require_scalar(evaluated)?); + } + values.push(output); + } + Ok(EvalValue::Array(FormulaArray::new(values).ok_or_else( + || { + calculation_error( + "use.office.spreadsheet_formula_array_invalid", + "Formula array constant is not rectangular.", + ) + }, + )?)) + } + } + } + + fn apply_binary( + &self, + operator: SpreadsheetFormulaBinaryOperator, + left: EvalValue, + right: EvalValue, + ) -> UseResult { + let left = self.materialize(left)?; + let right = self.materialize(right)?; + self.broadcast_binary(left, right, |left, right| { + scalar_binary(operator, left, right) + }) + } + + pub(super) fn map_numeric( + &self, + value: EvalValue, + operation: impl Fn(f64) -> ScalarValue + Copy, + ) -> UseResult { + let value = self.materialize(value)?; + Ok(match value { + EvalValue::Scalar(value) => EvalValue::Scalar(map_numeric_scalar(value, operation)), + EvalValue::Array(array) => EvalValue::Array(FormulaArray { + rows: array + .rows + .into_iter() + .map(|row| { + row.into_iter() + .map(|value| map_numeric_scalar(value, operation)) + .collect() + }) + .collect(), + }), + EvalValue::Reference(_) => { + return Err(calculation_error( + "use.office.spreadsheet_formula_calculation_invalid", + "Reference materialization did not produce a value.", + )); + } + }) + } + + pub(super) fn materialize(&self, value: EvalValue) -> UseResult { + match value { + EvalValue::Reference(areas) => self.materialize_areas(&areas), + value => Ok(value), + } + } + + pub(super) fn require_scalar(&self, value: EvalValue) -> UseResult { + match self.materialize(value)? { + EvalValue::Scalar(value) => Ok(value), + EvalValue::Array(array) if array.height() == 1 && array.width() == 1 => Ok(array + .rows + .into_iter() + .next() + .and_then(|row| row.into_iter().next()) + .unwrap_or(ScalarValue::Blank)), + EvalValue::Array(_) => Ok(ScalarValue::Error(SpreadsheetFormulaErrorLiteral::Value)), + EvalValue::Reference(_) => Ok(ScalarValue::Error( + SpreadsheetFormulaErrorLiteral::Reference, + )), + } + } + + pub(super) fn broadcast_binary( + &self, + left: EvalValue, + right: EvalValue, + operation: impl Fn(ScalarValue, ScalarValue) -> UseResult + Copy, + ) -> UseResult { + let left = into_array(left)?; + let right = into_array(right)?; + let height = broadcast_dimension(left.height(), right.height()); + let width = broadcast_dimension(left.width(), right.width()); + let (Some(height), Some(width)) = (height, width) else { + return Ok(EvalValue::Scalar(ScalarValue::Error( + SpreadsheetFormulaErrorLiteral::Value, + ))); + }; + checked_array_cells(height, width)?; + let mut rows = Vec::with_capacity(height); + let mut text_bytes = 0_usize; + for row in 0..height { + let mut values = Vec::with_capacity(width); + for column in 0..width { + let left_value = left + .broadcast_value(row, column) + .cloned() + .ok_or_else(invalid_array_shape)?; + let right_value = right + .broadcast_value(row, column) + .cloned() + .ok_or_else(invalid_array_shape)?; + let value = operation(left_value, right_value)?; + accumulate_formula_text_bytes(&mut text_bytes, &value)?; + values.push(value); + } + rows.push(values); + } + let array = FormulaArray { rows }; + if array.height() == 1 && array.width() == 1 { + Ok(EvalValue::Scalar( + array + .rows + .first() + .and_then(|row| row.first()) + .cloned() + .ok_or_else(invalid_array_shape)?, + )) + } else { + Ok(EvalValue::Array(array)) + } + } +} + +fn map_numeric_scalar(value: ScalarValue, operation: impl Fn(f64) -> ScalarValue) -> ScalarValue { + match scalar_number(value) { + Ok(number) => operation(number), + Err(error) => ScalarValue::Error(error), + } +} + +fn scalar_number(value: ScalarValue) -> Result { + match value { + ScalarValue::Blank => Ok(0.0), + ScalarValue::Number(value) => Ok(value), + ScalarValue::Boolean(value) => Ok(if value { 1.0 } else { 0.0 }), + ScalarValue::Text(value) => value + .parse::() + .ok() + .filter(|value| value.is_finite()) + .ok_or(SpreadsheetFormulaErrorLiteral::Value), + ScalarValue::Error(error) => Err(error), + } +} + +fn scalar_boolean(value: ScalarValue) -> Result { + match value { + ScalarValue::Blank => Ok(false), + ScalarValue::Number(value) => Ok(value != 0.0), + ScalarValue::Boolean(value) => Ok(value), + ScalarValue::Text(value) if value.eq_ignore_ascii_case("TRUE") => Ok(true), + ScalarValue::Text(value) if value.eq_ignore_ascii_case("FALSE") => Ok(false), + ScalarValue::Text(_) => Err(SpreadsheetFormulaErrorLiteral::Value), + ScalarValue::Error(error) => Err(error), + } +} + +fn scalar_text(value: ScalarValue) -> Result { + match value { + ScalarValue::Blank => Ok(String::new()), + ScalarValue::Number(value) => Ok(format_number(value)), + ScalarValue::Text(value) => Ok(value), + ScalarValue::Boolean(true) => Ok("TRUE".to_string()), + ScalarValue::Boolean(false) => Ok("FALSE".to_string()), + ScalarValue::Error(error) => Err(error), + } +} diff --git a/crates/office/src/spreadsheet_formula/evaluate/context.rs b/crates/office/src/spreadsheet_formula/evaluate/context.rs new file mode 100644 index 00000000..778ae78e --- /dev/null +++ b/crates/office/src/spreadsheet_formula/evaluate/context.rs @@ -0,0 +1,406 @@ +use std::collections::{BTreeMap, BTreeSet}; + +use a3s_use_core::{UseError, UseResult}; + +use crate::semantic::{DocumentNode, NativeOfficeDocument, OfficeNodeType}; +use crate::spreadsheet_reference::{CellRange, CellReference}; + +use super::super::{ + parse_spreadsheet_formula, SpreadsheetFormula, SpreadsheetFormulaCell, + SpreadsheetFormulaDependencyNode, SpreadsheetFormulaErrorLiteral, + SpreadsheetFormulaFunctionRegistry, +}; +use super::{calculation_error, spill_limit_error, MAX_SPREADSHEET_FORMULA_SPILL_CELLS}; +use crate::spreadsheet_formula::structured_reference::FormulaTableCatalog; + +#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord)] +pub(super) struct FormulaCellKey { + pub(super) sheet: usize, + pub(super) column: u32, + pub(super) row: u32, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub(super) struct FormulaReferenceArea { + pub(super) sheet: usize, + pub(super) start_column: u32, + pub(super) start_row: u32, + pub(super) end_column: u32, + pub(super) end_row: u32, +} + +impl FormulaReferenceArea { + pub(super) fn cell_count(self) -> Option { + let columns = u64::from(self.end_column - self.start_column + 1); + let rows = u64::from(self.end_row - self.start_row + 1); + usize::try_from(columns.checked_mul(rows)?).ok() + } + + pub(super) fn contains(self, key: FormulaCellKey) -> bool { + self.sheet == key.sheet + && (self.start_column..=self.end_column).contains(&key.column) + && (self.start_row..=self.end_row).contains(&key.row) + } + + pub(super) fn intersect(self, other: Self) -> Option { + if self.sheet != other.sheet { + return None; + } + let start_column = self.start_column.max(other.start_column); + let start_row = self.start_row.max(other.start_row); + let end_column = self.end_column.min(other.end_column); + let end_row = self.end_row.min(other.end_row); + (start_column <= end_column && start_row <= end_row).then_some(Self { + sheet: self.sheet, + start_column, + start_row, + end_column, + end_row, + }) + } +} + +#[derive(Debug, Clone, PartialEq)] +pub(super) enum ScalarValue { + Blank, + Number(f64), + Text(String), + Boolean(bool), + Error(SpreadsheetFormulaErrorLiteral), +} + +#[derive(Debug, Clone, PartialEq)] +pub(super) struct FormulaArray { + pub(super) rows: Vec>, +} + +impl FormulaArray { + pub(super) fn new(rows: Vec>) -> Option { + let width = rows.first()?.len(); + (width > 0 && rows.iter().all(|row| row.len() == width)).then_some(Self { rows }) + } + + pub(super) fn height(&self) -> usize { + self.rows.len() + } + + pub(super) fn width(&self) -> usize { + self.rows.first().map_or(0, Vec::len) + } + + pub(super) fn scalar(value: ScalarValue) -> Self { + Self { + rows: vec![vec![value]], + } + } + + pub(super) fn broadcast_value(&self, row: usize, column: usize) -> Option<&ScalarValue> { + let row = if self.height() == 1 { 0 } else { row }; + let column = if self.width() == 1 { 0 } else { column }; + self.rows.get(row)?.get(column) + } +} + +#[derive(Debug, Clone, PartialEq)] +pub(super) enum EvalValue { + Scalar(ScalarValue), + Array(FormulaArray), + Reference(Vec), +} + +#[derive(Debug, Clone)] +pub(super) struct NamedFormulaDefinition { + pub(super) name: String, + pub(super) formula: String, + pub(super) scope_sheet: Option, +} + +pub(super) struct FormulaRecord { + pub(super) key: FormulaCellKey, + pub(super) formula: SpreadsheetFormula, +} + +pub(super) struct EvaluationContext<'a> { + pub(super) sheet_names: Vec, + pub(super) registry: &'a SpreadsheetFormulaFunctionRegistry, + pub(super) values: BTreeMap, + pub(super) occupied: BTreeSet, + pub(super) formula_cells: BTreeSet, + pub(super) old_spills: BTreeMap, + pub(super) spills: BTreeMap, + pub(super) spill_owners: BTreeMap, + pub(super) named_definitions: BTreeMap<(Option, String), NamedFormulaDefinition>, + pub(super) named_stack: BTreeSet<(Option, String)>, + pub(super) tables: FormulaTableCatalog, +} + +pub(super) fn build_context<'a>( + document: &NativeOfficeDocument, + registry: &'a SpreadsheetFormulaFunctionRegistry, + graph_nodes: &[SpreadsheetFormulaDependencyNode], +) -> UseResult<(EvaluationContext<'a>, Vec)> { + let sheet_names = document + .root() + .children + .iter() + .filter(|node| node.node_type == OfficeNodeType::Worksheet) + .map(|node| node.path.trim_start_matches('/').to_string()) + .collect::>(); + let mut values = BTreeMap::new(); + let mut occupied = BTreeSet::new(); + let mut formula_cells = BTreeSet::new(); + let mut old_spills = BTreeMap::new(); + let mut old_spill_cell_count = 0_usize; + for (sheet, sheet_name) in sheet_names.iter().enumerate() { + let Some(sheet_node) = document.root().children.iter().find(|node| { + node.node_type == OfficeNodeType::Worksheet + && node + .path + .strip_prefix('/') + .is_some_and(|path| path.eq_ignore_ascii_case(sheet_name)) + }) else { + continue; + }; + for row in &sheet_node.children { + if row.node_type != OfficeNodeType::Row { + continue; + } + for cell in &row.children { + if cell.node_type != OfficeNodeType::Cell { + continue; + } + let reference = cell + .path + .rsplit_once('/') + .and_then(|(_, value)| CellReference::parse(value).ok()) + .ok_or_else(|| invalid_semantic_cell(cell))?; + let key = FormulaCellKey { + sheet, + column: reference.column, + row: reference.row, + }; + let is_formula = cell.format.contains_key("formula"); + let has_value = cell.format.get("valuePresent").map(String::as_str) == Some("true"); + if is_formula || has_value { + occupied.insert(key); + } + if is_formula { + formula_cells.insert(key); + if let Some(area) = stored_formula_spill(cell, key)? { + let spill_cells = area + .cell_count() + .ok_or_else(spill_limit_error)? + .saturating_sub(1); + old_spill_cell_count = old_spill_cell_count + .checked_add(spill_cells) + .ok_or_else(spill_limit_error)?; + if old_spill_cell_count > MAX_SPREADSHEET_FORMULA_SPILL_CELLS { + return Err( + spill_limit_error().with_detail("cells", old_spill_cell_count) + ); + } + old_spills.insert(key, area); + } + } else { + values.insert(key, semantic_scalar(cell)); + } + } + } + } + for (anchor, area) in &old_spills { + for row in area.start_row..=area.end_row { + for column in area.start_column..=area.end_column { + let key = FormulaCellKey { + sheet: area.sheet, + column, + row, + }; + if key != *anchor { + values.remove(&key); + } + } + } + } + let tables = FormulaTableCatalog::collect(document.root(), &sheet_names)?; + let named_definitions = collect_named_definitions(document.root(), &sheet_names); + let records = graph_nodes + .iter() + .map(|node| { + let key = public_cell_key(&node.cell, &sheet_names)?; + Ok(FormulaRecord { + key, + formula: parse_spreadsheet_formula(&node.formula) + .map_err(|error| error.with_detail("cell", node.cell.path()))?, + }) + }) + .collect::>>()?; + Ok(( + EvaluationContext { + sheet_names, + registry, + values, + occupied, + formula_cells, + old_spills, + spills: BTreeMap::new(), + spill_owners: BTreeMap::new(), + named_definitions, + named_stack: BTreeSet::new(), + tables, + }, + records, + )) +} + +pub(super) fn public_cell_key( + cell: &SpreadsheetFormulaCell, + sheet_names: &[String], +) -> UseResult { + let sheet = sheet_names + .iter() + .position(|sheet| sheet.eq_ignore_ascii_case(&cell.sheet)) + .ok_or_else(|| { + calculation_error( + "use.office.spreadsheet_formula_calculation_invalid", + format!("Calculation references missing worksheet '{}'.", cell.sheet), + ) + })?; + Ok(FormulaCellKey { + sheet, + column: cell.column, + row: cell.row, + }) +} + +fn semantic_scalar(cell: &DocumentNode) -> ScalarValue { + match cell.format.get("valueType").map(String::as_str) { + Some("String" | "Date") => ScalarValue::Text(cell.text.clone()), + Some("Boolean") => ScalarValue::Boolean(cell.text.eq_ignore_ascii_case("true")), + Some("Error") => SpreadsheetFormulaErrorLiteral::parse(&cell.text).map_or( + ScalarValue::Error(SpreadsheetFormulaErrorLiteral::Value), + ScalarValue::Error, + ), + _ if cell.text.is_empty() => ScalarValue::Blank, + _ => cell + .text + .parse::() + .ok() + .filter(|value| value.is_finite()) + .map_or( + ScalarValue::Error(SpreadsheetFormulaErrorLiteral::Value), + ScalarValue::Number, + ), + } +} + +fn collect_named_definitions( + root: &DocumentNode, + sheet_names: &[String], +) -> BTreeMap<(Option, String), NamedFormulaDefinition> { + root.children + .iter() + .filter(|node| node.node_type == OfficeNodeType::NamedRangeCollection) + .flat_map(|collection| &collection.children) + .filter_map(|node| { + let name = node.format.get("name")?.clone(); + let formula = node.format.get("ref")?.clone(); + let scope = node.format.get("scope")?; + let scope_sheet = if scope.eq_ignore_ascii_case("workbook") { + None + } else { + let sheet = scope.strip_prefix("worksheet:").unwrap_or(scope); + sheet_names + .iter() + .position(|candidate| candidate.eq_ignore_ascii_case(sheet)) + }; + Some(( + (scope_sheet, name.to_lowercase()), + NamedFormulaDefinition { + name, + formula, + scope_sheet, + }, + )) + }) + .collect() +} + +fn invalid_semantic_cell(cell: &DocumentNode) -> UseError { + calculation_error( + "use.office.spreadsheet_formula_calculation_invalid", + format!( + "Spreadsheet cell '{}' has an invalid coordinate.", + cell.path + ), + ) + .with_detail("cell", cell.path.clone()) +} + +fn stored_formula_spill( + cell: &DocumentNode, + anchor: FormulaCellKey, +) -> UseResult> { + let formula_type = cell.format.get("formulaType").map(String::as_str); + let formula_reference = cell.format.get("formulaRef"); + match (formula_type, formula_reference) { + (None, None) => Ok(None), + (Some(kind), None) if kind.eq_ignore_ascii_case("normal") => Ok(None), + (Some(kind), Some(reference)) if kind.eq_ignore_ascii_case("array") => { + let range = CellRange::parse(reference).map_err(|error| { + calculation_error( + "use.office.spreadsheet_formula_storage_invalid", + format!( + "Formula cell '{}' has invalid array range '{reference}': {error}", + cell.path + ), + ) + })?; + let anchor_reference = CellReference { + column: anchor.column, + row: anchor.row, + }; + if !range.contains(anchor_reference) { + return Err(calculation_error( + "use.office.spreadsheet_formula_storage_invalid", + format!( + "Formula array range '{}' does not contain anchor '{}'.", + range.a1(), + cell.path + ), + )); + } + Ok(Some(FormulaReferenceArea { + sheet: anchor.sheet, + start_column: range.start.column, + start_row: range.start.row, + end_column: range.end.column, + end_row: range.end.row, + })) + } + (Some(kind), None) if kind.eq_ignore_ascii_case("array") => Err(calculation_error( + "use.office.spreadsheet_formula_storage_invalid", + format!("Array formula cell '{}' has no array range.", cell.path), + )), + (None, Some(_)) => Err(calculation_error( + "use.office.spreadsheet_formula_storage_invalid", + format!( + "Formula cell '{}' has an array range without array storage type.", + cell.path + ), + )), + (Some(kind), Some(_)) if kind.eq_ignore_ascii_case("normal") => Err(calculation_error( + "use.office.spreadsheet_formula_storage_invalid", + format!( + "Normal formula cell '{}' cannot own an array range.", + cell.path + ), + )), + (Some(kind), _) => Err(calculation_error( + "use.office.spreadsheet_formula_storage_unsupported", + format!( + "Formula storage type '{kind}' at '{}' is not supported by native calculation.", + cell.path + ), + )), + } +} diff --git a/crates/office/src/spreadsheet_formula/evaluate/function.rs b/crates/office/src/spreadsheet_formula/evaluate/function.rs new file mode 100644 index 00000000..1fe53e92 --- /dev/null +++ b/crates/office/src/spreadsheet_formula/evaluate/function.rs @@ -0,0 +1,177 @@ +mod aggregate; +mod array; +mod lazy; + +use a3s_use_core::UseResult; + +use crate::spreadsheet_formula::registry::BuiltinFunction; +use crate::spreadsheet_formula::SpreadsheetFormulaExpression; +use crate::SpreadsheetFormulaErrorLiteral; + +use super::{ + calculation_error, checked_array_cells, finite_or_number_error, scalar_number, EvalValue, + EvaluationContext, FormulaCellKey, ScalarValue, +}; +use aggregate::{ + add_function_cells, aggregate, concatenate, count, count_a, logical_aggregate, logical_not, + Aggregate, +}; +use array::{row_or_column, sequence, transpose}; +use lazy::{evaluate_if, evaluate_if_error}; + +pub(super) fn evaluate_function( + context: &mut EvaluationContext<'_>, + name: &str, + arguments: &[Option], + current: FormulaCellKey, +) -> UseResult { + let Some(definition) = context.registry.get(name) else { + return Err(calculation_error( + "use.office.spreadsheet_formula_function_unsupported", + format!("Native calculation does not implement function '{name}'."), + ) + .with_detail("function", name)); + }; + if arguments.len() < definition.minimum_arguments + || definition + .maximum_arguments + .is_some_and(|maximum| arguments.len() > maximum) + { + return Err(calculation_error( + "use.office.spreadsheet_formula_function_arity", + format!( + "Function '{}' accepts {}{} arguments, but received {}.", + definition.name, + definition.minimum_arguments, + definition.maximum_arguments.map_or_else( + || " or more".to_string(), + |maximum| { + if maximum == definition.minimum_arguments { + String::new() + } else { + format!("-{maximum}") + } + } + ), + arguments.len() + ), + ) + .with_detail("function", definition.name.clone()) + .with_detail("arguments", arguments.len())); + } + let function = context.registry.function(name).ok_or_else(|| { + calculation_error( + "use.office.spreadsheet_formula_function_unsupported", + format!("Native function registry has no implementation for '{name}'."), + ) + })?; + if matches!(function, BuiltinFunction::If) { + return evaluate_if(context, arguments, current); + } + if matches!(function, BuiltinFunction::IfError) { + return evaluate_if_error(context, arguments, current); + } + let values = evaluate_arguments(context, arguments, current)?; + match function { + BuiltinFunction::Sum => aggregate(context, &values, Aggregate::Sum), + BuiltinFunction::Average => aggregate(context, &values, Aggregate::Average), + BuiltinFunction::Minimum => aggregate(context, &values, Aggregate::Minimum), + BuiltinFunction::Maximum => aggregate(context, &values, Aggregate::Maximum), + BuiltinFunction::Count => count(context, &values), + BuiltinFunction::CountA => count_a(context, &values), + BuiltinFunction::Absolute => context.map_numeric(argument(&values, 0)?, |number| { + finite_or_number_error(number.abs()) + }), + BuiltinFunction::SquareRoot => context.map_numeric(argument(&values, 0)?, |number| { + if number < 0.0 { + ScalarValue::Error(SpreadsheetFormulaErrorLiteral::Number) + } else { + finite_or_number_error(number.sqrt()) + } + }), + BuiltinFunction::Power => numeric_binary(context, &values, |left, right| { + finite_or_number_error(left.powf(right)) + }), + BuiltinFunction::Modulo => numeric_binary(context, &values, |left, right| { + if right == 0.0 { + ScalarValue::Error(SpreadsheetFormulaErrorLiteral::DivisionByZero) + } else { + finite_or_number_error(left - right * (left / right).floor()) + } + }), + BuiltinFunction::Round => numeric_binary(context, &values, |number, digits| { + if !(-308.0..=308.0).contains(&digits) { + return ScalarValue::Error(SpreadsheetFormulaErrorLiteral::Number); + } + let digits = digits.trunc() as i32; + let factor = 10_f64.powi(digits); + if !factor.is_finite() || factor == 0.0 { + return ScalarValue::Error(SpreadsheetFormulaErrorLiteral::Number); + } + finite_or_number_error((number * factor).round() / factor) + }), + BuiltinFunction::And => logical_aggregate(context, &values, true), + BuiltinFunction::Or => logical_aggregate(context, &values, false), + BuiltinFunction::Not => logical_not(context, argument(&values, 0)?), + BuiltinFunction::Concatenate => concatenate(context, &values), + BuiltinFunction::Row => row_or_column(context, &values, current, true), + BuiltinFunction::Column => row_or_column(context, &values, current, false), + BuiltinFunction::Sequence => sequence(context, &values), + BuiltinFunction::Transpose => transpose(context, &values), + BuiltinFunction::Pi => Ok(EvalValue::Scalar(ScalarValue::Number(std::f64::consts::PI))), + BuiltinFunction::NotAvailable => Ok(EvalValue::Scalar(ScalarValue::Error( + SpreadsheetFormulaErrorLiteral::NotAvailable, + ))), + BuiltinFunction::If | BuiltinFunction::IfError => Err(calculation_error( + "use.office.spreadsheet_formula_calculation_invalid", + "Lazy function reached eager dispatch.", + )), + } +} + +fn evaluate_arguments( + context: &mut EvaluationContext<'_>, + arguments: &[Option], + current: FormulaCellKey, +) -> UseResult> { + let mut values = Vec::with_capacity(arguments.len()); + let mut materialized_cells = 0_usize; + for argument in arguments { + let value = argument.as_ref().map_or_else( + || Ok(EvalValue::Scalar(ScalarValue::Blank)), + |argument| context.evaluate_expression(argument, current), + )?; + if let EvalValue::Array(array) = &value { + let cells = checked_array_cells(array.height(), array.width())?; + add_function_cells(&mut materialized_cells, cells)?; + } + values.push(value); + } + Ok(values) +} + +fn numeric_binary( + context: &EvaluationContext<'_>, + arguments: &[EvalValue], + operation: impl Fn(f64, f64) -> ScalarValue + Copy, +) -> UseResult { + let left = context.materialize(argument(arguments, 0)?)?; + let right = context.materialize(argument(arguments, 1)?)?; + context.broadcast_binary(left, right, |left, right| { + let (Ok(left), Ok(right)) = (scalar_number(left), scalar_number(right)) else { + return Ok(ScalarValue::Error(SpreadsheetFormulaErrorLiteral::Value)); + }; + Ok(operation(left, right)) + }) +} + +fn argument(arguments: &[EvalValue], index: usize) -> UseResult { + arguments.get(index).cloned().ok_or_else(missing_argument) +} + +fn missing_argument() -> a3s_use_core::UseError { + calculation_error( + "use.office.spreadsheet_formula_calculation_invalid", + "Native formula function dispatch is missing a validated argument.", + ) +} diff --git a/crates/office/src/spreadsheet_formula/evaluate/function/aggregate.rs b/crates/office/src/spreadsheet_formula/evaluate/function/aggregate.rs new file mode 100644 index 00000000..3258e293 --- /dev/null +++ b/crates/office/src/spreadsheet_formula/evaluate/function/aggregate.rs @@ -0,0 +1,392 @@ +use a3s_use_core::UseResult; + +use crate::SpreadsheetFormulaErrorLiteral; + +use super::super::{ + calculation_error, scalar_boolean, scalar_number, scalar_text, EvalValue, EvaluationContext, + FormulaArray, ScalarValue, MAX_SPREADSHEET_FORMULA_SPILL_CELLS, +}; + +#[derive(Debug, Clone, Copy)] +pub(super) enum Aggregate { + Sum, + Average, + Minimum, + Maximum, +} + +pub(super) fn aggregate( + context: &EvaluationContext<'_>, + arguments: &[EvalValue], + operation: Aggregate, +) -> UseResult { + let mut count = 0_usize; + let mut sum = 0.0_f64; + let mut minimum = None; + let mut maximum = None; + let mut materialized_cells = 0_usize; + for argument in arguments { + match argument { + EvalValue::Reference(areas) => { + let values = context.populated_reference_values(areas)?; + add_function_cells(&mut materialized_cells, values.len())?; + for value in values { + match value { + ScalarValue::Number(value) => { + record_number(&mut count, &mut sum, &mut minimum, &mut maximum, value) + } + ScalarValue::Error(error) => { + return Ok(EvalValue::Scalar(ScalarValue::Error(error))); + } + _ => {} + } + } + } + EvalValue::Array(array) => { + add_function_cells( + &mut materialized_cells, + array.height().saturating_mul(array.width()), + )?; + for value in array.rows.iter().flatten().cloned() { + match scalar_number(value) { + Ok(value) => { + record_number(&mut count, &mut sum, &mut minimum, &mut maximum, value) + } + Err(error) => { + return Ok(EvalValue::Scalar(ScalarValue::Error(error))); + } + } + } + } + EvalValue::Scalar(value) => { + add_function_cells(&mut materialized_cells, 1)?; + match scalar_number(value.clone()) { + Ok(value) => { + record_number(&mut count, &mut sum, &mut minimum, &mut maximum, value) + } + Err(error) => return Ok(EvalValue::Scalar(ScalarValue::Error(error))), + } + } + } + } + let value = match operation { + Aggregate::Sum => sum, + Aggregate::Average if count == 0 => { + return Ok(EvalValue::Scalar(ScalarValue::Error( + SpreadsheetFormulaErrorLiteral::DivisionByZero, + ))); + } + Aggregate::Average => sum / count as f64, + Aggregate::Minimum => minimum.unwrap_or(0.0), + Aggregate::Maximum => maximum.unwrap_or(0.0), + }; + Ok(EvalValue::Scalar(super::super::finite_or_number_error( + value, + ))) +} + +pub(super) fn count( + context: &EvaluationContext<'_>, + arguments: &[EvalValue], +) -> UseResult { + let mut count = 0_usize; + let mut materialized_cells = 0_usize; + for argument in arguments { + match argument { + EvalValue::Reference(areas) => { + let values = context.populated_reference_values(areas)?; + add_function_cells(&mut materialized_cells, values.len())?; + count = count.saturating_add( + values + .iter() + .filter(|value| matches!(value, ScalarValue::Number(_))) + .count(), + ); + } + EvalValue::Array(array) => { + add_function_cells( + &mut materialized_cells, + array.height().saturating_mul(array.width()), + )?; + count = count.saturating_add( + array + .rows + .iter() + .flatten() + .filter(|value| matches!(value, ScalarValue::Number(_))) + .count(), + ); + } + EvalValue::Scalar(value) => { + add_function_cells(&mut materialized_cells, 1)?; + match value { + ScalarValue::Number(_) | ScalarValue::Boolean(_) => { + count = count.saturating_add(1); + } + ScalarValue::Text(value) + if value + .parse::() + .ok() + .is_some_and(|value| value.is_finite()) => + { + count = count.saturating_add(1); + } + ScalarValue::Error(error) => { + return Ok(EvalValue::Scalar(ScalarValue::Error(*error))); + } + ScalarValue::Blank | ScalarValue::Text(_) => {} + } + } + } + } + Ok(EvalValue::Scalar(ScalarValue::Number(count as f64))) +} + +pub(super) fn count_a( + context: &EvaluationContext<'_>, + arguments: &[EvalValue], +) -> UseResult { + let mut count = 0_usize; + let mut materialized_cells = 0_usize; + for argument in arguments { + match argument { + EvalValue::Reference(areas) => { + let values = context.populated_reference_values(areas)?; + add_function_cells(&mut materialized_cells, values.len())?; + count = count.saturating_add( + values + .iter() + .filter(|value| !matches!(value, ScalarValue::Blank)) + .count(), + ); + } + EvalValue::Array(array) => { + add_function_cells( + &mut materialized_cells, + array.height().saturating_mul(array.width()), + )?; + count = count.saturating_add( + array + .rows + .iter() + .flatten() + .filter(|value| !matches!(value, ScalarValue::Blank)) + .count(), + ); + } + EvalValue::Scalar(ScalarValue::Blank) => { + add_function_cells(&mut materialized_cells, 1)?; + } + EvalValue::Scalar(_) => { + add_function_cells(&mut materialized_cells, 1)?; + count = count.saturating_add(1); + } + } + } + Ok(EvalValue::Scalar(ScalarValue::Number(count as f64))) +} + +pub(super) fn logical_aggregate( + context: &EvaluationContext<'_>, + arguments: &[EvalValue], + and: bool, +) -> UseResult { + let mut observed = false; + let mut result = and; + let mut materialized_cells = 0_usize; + for argument in arguments { + match argument { + EvalValue::Reference(areas) => { + let values = context.populated_reference_values(areas)?; + add_function_cells(&mut materialized_cells, values.len())?; + for value in values { + match value { + ScalarValue::Blank | ScalarValue::Text(_) => {} + ScalarValue::Error(error) => { + return Ok(EvalValue::Scalar(ScalarValue::Error(error))); + } + value => { + observed = true; + match scalar_boolean(value) { + Ok(value) => update_logical_result(&mut result, and, value), + Err(error) => { + return Ok(EvalValue::Scalar(ScalarValue::Error(error))); + } + } + } + } + } + } + EvalValue::Array(array) => { + add_function_cells( + &mut materialized_cells, + array.height().saturating_mul(array.width()), + )?; + for value in array.rows.iter().flatten().cloned() { + match value { + ScalarValue::Blank | ScalarValue::Text(_) => {} + ScalarValue::Error(error) => { + return Ok(EvalValue::Scalar(ScalarValue::Error(error))); + } + value => { + observed = true; + match scalar_boolean(value) { + Ok(value) => update_logical_result(&mut result, and, value), + Err(error) => { + return Ok(EvalValue::Scalar(ScalarValue::Error(error))); + } + } + } + } + } + } + EvalValue::Scalar(value) => { + add_function_cells(&mut materialized_cells, 1)?; + match scalar_boolean(value.clone()) { + Ok(value) => { + observed = true; + update_logical_result(&mut result, and, value); + } + Err(error) => { + return Ok(EvalValue::Scalar(ScalarValue::Error(error))); + } + } + } + } + } + if !observed { + return Ok(EvalValue::Scalar(ScalarValue::Error( + SpreadsheetFormulaErrorLiteral::Value, + ))); + } + Ok(EvalValue::Scalar(ScalarValue::Boolean(result))) +} + +pub(super) fn concatenate( + context: &EvaluationContext<'_>, + arguments: &[EvalValue], +) -> UseResult { + let mut output = String::new(); + let mut materialized_cells = 0_usize; + for argument in arguments { + let values = argument_scalars(context, argument)?; + add_function_cells(&mut materialized_cells, values.len())?; + for value in values { + match scalar_text(value) { + Ok(value) => append_formula_text(&mut output, &value)?, + Err(error) => return Ok(EvalValue::Scalar(ScalarValue::Error(error))), + } + } + } + Ok(EvalValue::Scalar(ScalarValue::Text(output))) +} + +pub(super) fn logical_not( + context: &EvaluationContext<'_>, + value: EvalValue, +) -> UseResult { + let value = context.materialize(value)?; + Ok(match value { + EvalValue::Scalar(value) => EvalValue::Scalar(not_scalar(value)), + EvalValue::Array(array) => EvalValue::Array(FormulaArray { + rows: array + .rows + .into_iter() + .map(|row| row.into_iter().map(not_scalar).collect()) + .collect(), + }), + EvalValue::Reference(_) => { + return Err(calculation_error( + "use.office.spreadsheet_formula_calculation_invalid", + "Reference materialization did not produce a logical value.", + )); + } + }) +} + +pub(super) fn add_function_cells(total: &mut usize, added: usize) -> UseResult<()> { + *total = total + .checked_add(added) + .ok_or_else(function_argument_limit)?; + if *total > MAX_SPREADSHEET_FORMULA_SPILL_CELLS { + return Err(function_argument_limit().with_detail("cells", *total)); + } + Ok(()) +} + +fn record_number( + count: &mut usize, + sum: &mut f64, + minimum: &mut Option, + maximum: &mut Option, + value: f64, +) { + *count = count.saturating_add(1); + *sum += value; + *minimum = Some(minimum.map_or(value, |current| current.min(value))); + *maximum = Some(maximum.map_or(value, |current| current.max(value))); +} + +fn function_argument_limit() -> a3s_use_core::UseError { + calculation_error( + "use.office.spreadsheet_formula_function_array_limit", + format!( + "Native formula functions materialize at most {MAX_SPREADSHEET_FORMULA_SPILL_CELLS} argument cells per call." + ), + ) +} + +fn update_logical_result(result: &mut bool, and: bool, value: bool) { + if and { + *result &= value; + } else { + *result |= value; + } +} + +fn not_scalar(value: ScalarValue) -> ScalarValue { + match scalar_boolean(value) { + Ok(value) => ScalarValue::Boolean(!value), + Err(error) => ScalarValue::Error(error), + } +} + +fn argument_scalars( + context: &EvaluationContext<'_>, + argument: &EvalValue, +) -> UseResult> { + match argument { + EvalValue::Scalar(value) => Ok(vec![value.clone()]), + EvalValue::Array(array) => Ok(array.rows.iter().flatten().cloned().collect()), + EvalValue::Reference(areas) => { + let value = context.materialize(EvalValue::Reference(areas.clone()))?; + match value { + EvalValue::Scalar(value) => Ok(vec![value]), + EvalValue::Array(array) => Ok(array.rows.into_iter().flatten().collect()), + EvalValue::Reference(_) => Ok(Vec::new()), + } + } + } +} + +fn append_formula_text(output: &mut String, value: &str) -> UseResult<()> { + let bytes = output + .len() + .checked_add(value.len()) + .ok_or_else(formula_text_limit)?; + if bytes > super::super::MAX_SPREADSHEET_FORMULA_TEXT_BYTES { + return Err(formula_text_limit().with_detail("bytes", bytes)); + } + output.push_str(value); + Ok(()) +} + +fn formula_text_limit() -> a3s_use_core::UseError { + calculation_error( + "use.office.spreadsheet_formula_text_limit", + format!( + "Native formula text results support at most {} UTF-8 bytes.", + super::super::MAX_SPREADSHEET_FORMULA_TEXT_BYTES + ), + ) +} diff --git a/crates/office/src/spreadsheet_formula/evaluate/function/array.rs b/crates/office/src/spreadsheet_formula/evaluate/function/array.rs new file mode 100644 index 00000000..000860d9 --- /dev/null +++ b/crates/office/src/spreadsheet_formula/evaluate/function/array.rs @@ -0,0 +1,145 @@ +use a3s_use_core::UseResult; + +use crate::SpreadsheetFormulaErrorLiteral; + +use super::super::{ + calculation_error, finite_or_number_error, into_array, invalid_array_shape, scalar_number, + EvalValue, EvaluationContext, FormulaArray, FormulaCellKey, ScalarValue, + MAX_SPREADSHEET_FORMULA_SPILL_CELLS, +}; +use super::{argument, missing_argument}; + +pub(super) fn row_or_column( + _context: &EvaluationContext<'_>, + arguments: &[EvalValue], + current: FormulaCellKey, + row: bool, +) -> UseResult { + if arguments.is_empty() { + return Ok(EvalValue::Scalar(ScalarValue::Number(f64::from(if row { + current.row + } else { + current.column + })))); + } + let EvalValue::Reference(areas) = arguments.first().ok_or_else(missing_argument)? else { + return Ok(EvalValue::Scalar(ScalarValue::Error( + SpreadsheetFormulaErrorLiteral::Value, + ))); + }; + let mut output = Vec::new(); + for area in areas { + let cells = area.cell_count().ok_or_else(function_array_limit)?; + if cells > MAX_SPREADSHEET_FORMULA_SPILL_CELLS + || output.len().saturating_add(cells) > MAX_SPREADSHEET_FORMULA_SPILL_CELLS + { + return Err(function_array_limit()); + } + for source_row in area.start_row..=area.end_row { + let mut values = Vec::new(); + for source_column in area.start_column..=area.end_column { + values.push(ScalarValue::Number(f64::from(if row { + source_row + } else { + source_column + }))); + } + output.push(values); + } + } + collapse_array(FormulaArray::new(output).ok_or_else(function_array_limit)?) +} + +pub(super) fn sequence( + context: &EvaluationContext<'_>, + arguments: &[EvalValue], +) -> UseResult { + let rows = positive_dimension(context.require_scalar(argument(arguments, 0)?)?); + let columns = if let Some(value) = arguments.get(1) { + positive_dimension(context.require_scalar(value.clone())?) + } else { + Ok(1) + }; + let start = if let Some(value) = arguments.get(2) { + scalar_number(context.require_scalar(value.clone())?) + } else { + Ok(1.0) + }; + let step = if let Some(value) = arguments.get(3) { + scalar_number(context.require_scalar(value.clone())?) + } else { + Ok(1.0) + }; + let (Ok(rows), Ok(columns), Ok(start), Ok(step)) = (rows, columns, start, step) else { + return Ok(EvalValue::Scalar(ScalarValue::Error( + SpreadsheetFormulaErrorLiteral::Value, + ))); + }; + let cells = rows.checked_mul(columns).ok_or_else(function_array_limit)?; + if cells > MAX_SPREADSHEET_FORMULA_SPILL_CELLS { + return Err(function_array_limit().with_detail("cells", cells)); + } + let mut output = Vec::with_capacity(rows); + for row in 0..rows { + let mut values = Vec::with_capacity(columns); + for column in 0..columns { + let offset = row + .checked_mul(columns) + .and_then(|value| value.checked_add(column)) + .ok_or_else(function_array_limit)?; + values.push(finite_or_number_error(start + step * offset as f64)); + } + output.push(values); + } + Ok(EvalValue::Array(FormulaArray { rows: output })) +} + +pub(super) fn transpose( + context: &EvaluationContext<'_>, + arguments: &[EvalValue], +) -> UseResult { + let array = into_array(context.materialize(argument(arguments, 0)?)?)?; + let mut output = vec![vec![ScalarValue::Blank; array.height()]; array.width()]; + for (row, values) in array.rows.into_iter().enumerate() { + for (column, value) in values.into_iter().enumerate() { + let target = output + .get_mut(column) + .and_then(|values| values.get_mut(row)) + .ok_or_else(invalid_array_shape)?; + *target = value; + } + } + collapse_array(FormulaArray { rows: output }) +} + +pub(super) fn collapse_array(array: FormulaArray) -> UseResult { + if array.height() == 1 && array.width() == 1 { + Ok(EvalValue::Scalar( + array + .rows + .into_iter() + .next() + .and_then(|row| row.into_iter().next()) + .unwrap_or(ScalarValue::Blank), + )) + } else { + Ok(EvalValue::Array(array)) + } +} + +fn positive_dimension(value: ScalarValue) -> Result { + let value = scalar_number(value)?; + if !value.is_finite() || value < 1.0 || value.fract() != 0.0 { + return Err(SpreadsheetFormulaErrorLiteral::Value); + } + usize::try_from(value as u64).map_err(|_| SpreadsheetFormulaErrorLiteral::Number) +} + +fn function_array_limit() -> a3s_use_core::UseError { + calculation_error( + "use.office.spreadsheet_formula_function_array_limit", + format!( + "Native formula functions return at most {MAX_SPREADSHEET_FORMULA_SPILL_CELLS} array cells." + ), + ) +} diff --git a/crates/office/src/spreadsheet_formula/evaluate/function/lazy.rs b/crates/office/src/spreadsheet_formula/evaluate/function/lazy.rs new file mode 100644 index 00000000..18476a4b --- /dev/null +++ b/crates/office/src/spreadsheet_formula/evaluate/function/lazy.rs @@ -0,0 +1,170 @@ +use a3s_use_core::UseResult; + +use crate::spreadsheet_formula::SpreadsheetFormulaExpression; +use crate::SpreadsheetFormulaErrorLiteral; + +use super::super::{ + accumulate_formula_text_bytes, broadcast_dimension, checked_array_cells, into_array, + invalid_array_shape, scalar_boolean, EvalValue, EvaluationContext, FormulaArray, + FormulaCellKey, ScalarValue, +}; +use super::array::collapse_array; + +pub(super) fn evaluate_if( + context: &mut EvaluationContext<'_>, + arguments: &[Option], + current: FormulaCellKey, +) -> UseResult { + let condition = + evaluate_optional_argument(context, arguments.first(), current, ScalarValue::Blank)?; + let condition = context.materialize(condition)?; + if let EvalValue::Scalar(condition) = condition { + return match scalar_boolean(condition) { + Ok(true) => { + evaluate_optional_argument(context, arguments.get(1), current, ScalarValue::Blank) + } + Ok(false) => evaluate_optional_argument( + context, + arguments.get(2), + current, + ScalarValue::Boolean(false), + ), + Err(error) => Ok(EvalValue::Scalar(ScalarValue::Error(error))), + }; + } + let true_value = + evaluate_optional_argument(context, arguments.get(1), current, ScalarValue::Blank)?; + let false_value = evaluate_optional_argument( + context, + arguments.get(2), + current, + ScalarValue::Boolean(false), + )?; + select_array(context, condition, true_value, false_value) +} + +pub(super) fn evaluate_if_error( + context: &mut EvaluationContext<'_>, + arguments: &[Option], + current: FormulaCellKey, +) -> UseResult { + let value = + evaluate_optional_argument(context, arguments.first(), current, ScalarValue::Blank)?; + let value = context.materialize(value)?; + if let EvalValue::Scalar(ScalarValue::Error(_)) = value { + return evaluate_optional_argument(context, arguments.get(1), current, ScalarValue::Blank); + } + let EvalValue::Array(array) = &value else { + return Ok(value); + }; + if !array + .rows + .iter() + .flatten() + .any(|value| matches!(value, ScalarValue::Error(_))) + { + return Ok(value); + } + let fallback = + evaluate_optional_argument(context, arguments.get(1), current, ScalarValue::Blank)?; + let fallback = into_array(context.materialize(fallback)?)?; + let height = broadcast_dimension(array.height(), fallback.height()); + let width = broadcast_dimension(array.width(), fallback.width()); + let (Some(height), Some(width)) = (height, width) else { + return Ok(EvalValue::Scalar(ScalarValue::Error( + SpreadsheetFormulaErrorLiteral::Value, + ))); + }; + checked_array_cells(height, width)?; + let mut rows = Vec::with_capacity(height); + let mut text_bytes = 0_usize; + for row in 0..height { + let mut values = Vec::with_capacity(width); + for column in 0..width { + let original = array + .broadcast_value(row, column) + .cloned() + .ok_or_else(invalid_array_shape)?; + let value = if matches!(original, ScalarValue::Error(_)) { + fallback + .broadcast_value(row, column) + .cloned() + .ok_or_else(invalid_array_shape)? + } else { + original + }; + accumulate_formula_text_bytes(&mut text_bytes, &value)?; + values.push(value); + } + rows.push(values); + } + collapse_array(FormulaArray { rows }) +} + +fn evaluate_optional_argument( + context: &mut EvaluationContext<'_>, + argument: Option<&Option>, + current: FormulaCellKey, + absent: ScalarValue, +) -> UseResult { + match argument { + Some(Some(argument)) => context.evaluate_expression(argument, current), + Some(None) => Ok(EvalValue::Scalar(ScalarValue::Blank)), + None => Ok(EvalValue::Scalar(absent)), + } +} + +fn select_array( + context: &EvaluationContext<'_>, + condition: EvalValue, + true_value: EvalValue, + false_value: EvalValue, +) -> UseResult { + let condition = into_array(condition)?; + let true_value = into_array(context.materialize(true_value)?)?; + let false_value = into_array(context.materialize(false_value)?)?; + let height = [ + condition.height(), + true_value.height(), + false_value.height(), + ] + .into_iter() + .reduce(|left, right| broadcast_dimension(left, right).unwrap_or(usize::MAX)) + .filter(|value| *value != usize::MAX); + let width = [condition.width(), true_value.width(), false_value.width()] + .into_iter() + .reduce(|left, right| broadcast_dimension(left, right).unwrap_or(usize::MAX)) + .filter(|value| *value != usize::MAX); + let (Some(height), Some(width)) = (height, width) else { + return Ok(EvalValue::Scalar(ScalarValue::Error( + SpreadsheetFormulaErrorLiteral::Value, + ))); + }; + checked_array_cells(height, width)?; + let mut rows = Vec::with_capacity(height); + let mut text_bytes = 0_usize; + for row in 0..height { + let mut values = Vec::with_capacity(width); + for column in 0..width { + let condition = condition + .broadcast_value(row, column) + .cloned() + .ok_or_else(invalid_array_shape)?; + let value = match scalar_boolean(condition) { + Ok(true) => true_value + .broadcast_value(row, column) + .cloned() + .ok_or_else(invalid_array_shape)?, + Ok(false) => false_value + .broadcast_value(row, column) + .cloned() + .ok_or_else(invalid_array_shape)?, + Err(error) => ScalarValue::Error(error), + }; + accumulate_formula_text_bytes(&mut text_bytes, &value)?; + values.push(value); + } + rows.push(values); + } + collapse_array(FormulaArray { rows }) +} diff --git a/crates/office/src/spreadsheet_formula/evaluate/operators.rs b/crates/office/src/spreadsheet_formula/evaluate/operators.rs new file mode 100644 index 00000000..6ed5b806 --- /dev/null +++ b/crates/office/src/spreadsheet_formula/evaluate/operators.rs @@ -0,0 +1,243 @@ +use a3s_use_core::{UseError, UseResult}; + +use crate::discovery::office_error; +use crate::{ + SpreadsheetFormulaBinaryOperator, SpreadsheetFormulaErrorLiteral, SpreadsheetFormulaValue, +}; + +use super::{ + scalar_number, scalar_text, EvalValue, FormulaArray, ScalarValue, + MAX_SPREADSHEET_FORMULA_CALCULATION_TEXT_BYTES, MAX_SPREADSHEET_FORMULA_SPILL_CELLS, + MAX_SPREADSHEET_FORMULA_TEXT_BYTES, +}; + +pub(super) fn scalar_binary( + operator: SpreadsheetFormulaBinaryOperator, + left: ScalarValue, + right: ScalarValue, +) -> UseResult { + if let ScalarValue::Error(error) = &left { + return Ok(ScalarValue::Error(*error)); + } + if let ScalarValue::Error(error) = &right { + return Ok(ScalarValue::Error(*error)); + } + Ok(match operator { + SpreadsheetFormulaBinaryOperator::Add + | SpreadsheetFormulaBinaryOperator::Subtract + | SpreadsheetFormulaBinaryOperator::Multiply + | SpreadsheetFormulaBinaryOperator::Divide + | SpreadsheetFormulaBinaryOperator::Power => { + let (Ok(left), Ok(right)) = (scalar_number(left), scalar_number(right)) else { + return Ok(ScalarValue::Error(SpreadsheetFormulaErrorLiteral::Value)); + }; + let value = match operator { + SpreadsheetFormulaBinaryOperator::Add => left + right, + SpreadsheetFormulaBinaryOperator::Subtract => left - right, + SpreadsheetFormulaBinaryOperator::Multiply => left * right, + SpreadsheetFormulaBinaryOperator::Divide if right == 0.0 => { + return Ok(ScalarValue::Error( + SpreadsheetFormulaErrorLiteral::DivisionByZero, + )); + } + SpreadsheetFormulaBinaryOperator::Divide => left / right, + SpreadsheetFormulaBinaryOperator::Power => left.powf(right), + _ => return Ok(ScalarValue::Error(SpreadsheetFormulaErrorLiteral::Value)), + }; + finite_or_number_error(value) + } + SpreadsheetFormulaBinaryOperator::Concatenate => { + match (scalar_text(left), scalar_text(right)) { + (Ok(left), Ok(right)) => { + let bytes = left + .len() + .checked_add(right.len()) + .ok_or_else(formula_text_limit_error)?; + if bytes > MAX_SPREADSHEET_FORMULA_TEXT_BYTES { + return Err(formula_text_limit_error().with_detail("bytes", bytes)); + } + let mut value = String::with_capacity(bytes); + value.push_str(&left); + value.push_str(&right); + ScalarValue::Text(value) + } + (Err(error), _) | (_, Err(error)) => ScalarValue::Error(error), + } + } + SpreadsheetFormulaBinaryOperator::Equal + | SpreadsheetFormulaBinaryOperator::NotEqual + | SpreadsheetFormulaBinaryOperator::LessThan + | SpreadsheetFormulaBinaryOperator::LessThanOrEqual + | SpreadsheetFormulaBinaryOperator::GreaterThan + | SpreadsheetFormulaBinaryOperator::GreaterThanOrEqual => { + let ordering = compare_scalars(&left, &right); + let value = match operator { + SpreadsheetFormulaBinaryOperator::Equal => ordering == std::cmp::Ordering::Equal, + SpreadsheetFormulaBinaryOperator::NotEqual => ordering != std::cmp::Ordering::Equal, + SpreadsheetFormulaBinaryOperator::LessThan => ordering.is_lt(), + SpreadsheetFormulaBinaryOperator::LessThanOrEqual => !ordering.is_gt(), + SpreadsheetFormulaBinaryOperator::GreaterThan => ordering.is_gt(), + SpreadsheetFormulaBinaryOperator::GreaterThanOrEqual => !ordering.is_lt(), + _ => false, + }; + ScalarValue::Boolean(value) + } + _ => ScalarValue::Error(SpreadsheetFormulaErrorLiteral::Value), + }) +} + +fn compare_scalars(left: &ScalarValue, right: &ScalarValue) -> std::cmp::Ordering { + match (left, right) { + (ScalarValue::Blank, ScalarValue::Blank) => std::cmp::Ordering::Equal, + (ScalarValue::Number(left), ScalarValue::Number(right)) => { + left.partial_cmp(right).unwrap_or(std::cmp::Ordering::Equal) + } + (ScalarValue::Boolean(left), ScalarValue::Boolean(right)) => left.cmp(right), + (ScalarValue::Text(left), ScalarValue::Text(right)) => { + left.to_lowercase().cmp(&right.to_lowercase()) + } + (ScalarValue::Blank, ScalarValue::Number(value)) + | (ScalarValue::Number(value), ScalarValue::Blank) + if *value == 0.0 => + { + std::cmp::Ordering::Equal + } + _ => scalar_rank(left).cmp(&scalar_rank(right)), + } +} + +fn scalar_rank(value: &ScalarValue) -> u8 { + match value { + ScalarValue::Blank => 0, + ScalarValue::Number(_) => 1, + ScalarValue::Text(_) => 2, + ScalarValue::Boolean(_) => 3, + ScalarValue::Error(_) => 4, + } +} + +pub(super) fn finite_or_number_error(value: f64) -> ScalarValue { + if value.is_finite() { + ScalarValue::Number(if value == 0.0 { 0.0 } else { value }) + } else { + ScalarValue::Error(SpreadsheetFormulaErrorLiteral::Number) + } +} + +pub(super) fn into_array(value: EvalValue) -> UseResult { + match value { + EvalValue::Scalar(value) => Ok(FormulaArray::scalar(value)), + EvalValue::Array(array) => Ok(array), + EvalValue::Reference(_) => Err(calculation_error( + "use.office.spreadsheet_formula_calculation_invalid", + "Reference must be materialized before array broadcasting.", + )), + } +} + +pub(super) fn broadcast_dimension(left: usize, right: usize) -> Option { + if left == right { + Some(left) + } else if left == 1 { + Some(right) + } else if right == 1 { + Some(left) + } else { + None + } +} + +pub(super) fn checked_array_cells(height: usize, width: usize) -> UseResult { + let cells = height.checked_mul(width).ok_or_else(spill_limit_error)?; + if cells > MAX_SPREADSHEET_FORMULA_SPILL_CELLS { + return Err(spill_limit_error().with_detail("cells", cells)); + } + Ok(cells) +} + +pub(super) fn public_scalar(value: &ScalarValue) -> SpreadsheetFormulaValue { + match value { + ScalarValue::Blank => SpreadsheetFormulaValue::Blank, + ScalarValue::Number(value) => SpreadsheetFormulaValue::Number { + value: format_number(*value), + }, + ScalarValue::Text(value) => SpreadsheetFormulaValue::Text { + value: value.clone(), + }, + ScalarValue::Boolean(value) => SpreadsheetFormulaValue::Boolean { value: *value }, + ScalarValue::Error(error) => SpreadsheetFormulaValue::Error { error: *error }, + } +} + +pub(super) fn ensure_formula_text_limit(value: &ScalarValue) -> UseResult<()> { + let ScalarValue::Text(value) = value else { + return Ok(()); + }; + if value.len() > MAX_SPREADSHEET_FORMULA_TEXT_BYTES { + return Err(formula_text_limit_error().with_detail("bytes", value.len())); + } + Ok(()) +} + +pub(super) fn accumulate_formula_text_bytes( + bytes: &mut usize, + value: &ScalarValue, +) -> UseResult<()> { + ensure_formula_text_limit(value)?; + let ScalarValue::Text(value) = value else { + return Ok(()); + }; + *bytes = bytes + .checked_add(value.len()) + .ok_or_else(calculation_text_limit_error)?; + if *bytes > MAX_SPREADSHEET_FORMULA_CALCULATION_TEXT_BYTES { + return Err(calculation_text_limit_error().with_detail("bytes", *bytes)); + } + Ok(()) +} + +pub(super) fn format_number(value: f64) -> String { + if value == 0.0 { + "0".to_string() + } else { + value.to_string() + } +} + +pub(super) fn spill_limit_error() -> UseError { + calculation_error( + "use.office.spreadsheet_formula_spill_limit", + format!( + "Native formula spills support at most {MAX_SPREADSHEET_FORMULA_SPILL_CELLS} cells within worksheet limits." + ), + ) +} + +fn formula_text_limit_error() -> UseError { + calculation_error( + "use.office.spreadsheet_formula_text_limit", + format!( + "Native formula text results support at most {MAX_SPREADSHEET_FORMULA_TEXT_BYTES} UTF-8 bytes." + ), + ) +} + +pub(super) fn calculation_text_limit_error() -> UseError { + calculation_error( + "use.office.spreadsheet_formula_text_limit", + format!( + "Native formula calculation produces at most {MAX_SPREADSHEET_FORMULA_CALCULATION_TEXT_BYTES} cumulative UTF-8 text-result bytes." + ), + ) +} + +pub(super) fn invalid_array_shape() -> UseError { + calculation_error( + "use.office.spreadsheet_formula_calculation_invalid", + "Native formula calculation produced an invalid array shape.", + ) +} + +pub(super) fn calculation_error(code: &str, message: impl Into) -> UseError { + office_error(code, message) +} diff --git a/crates/office/src/spreadsheet_formula/evaluate/reference.rs b/crates/office/src/spreadsheet_formula/evaluate/reference.rs new file mode 100644 index 00000000..d0344b37 --- /dev/null +++ b/crates/office/src/spreadsheet_formula/evaluate/reference.rs @@ -0,0 +1,615 @@ +use a3s_use_core::UseResult; + +use crate::spreadsheet_formula::structured_reference::StructuredReferenceErrorKind; +use crate::spreadsheet_formula::{ + parse_spreadsheet_formula, SpreadsheetFormulaBinaryOperator, SpreadsheetFormulaExpression, + SpreadsheetFormulaExpressionKind, SpreadsheetFormulaQualifier, SpreadsheetFormulaReference, + SpreadsheetFormulaReferenceKind, MAX_SPREADSHEET_FORMULA_DEPTH, + MAX_SPREADSHEET_FORMULA_REFERENCE_AREAS, MAX_SPREADSHEET_FORMULA_REFERENCE_VISITS, +}; + +use super::{ + accumulate_formula_text_bytes, calculation_error, invalid_array_shape, EvalValue, + EvaluationContext, FormulaArray, FormulaCellKey, FormulaReferenceArea, ScalarValue, + MAX_SPREADSHEET_FORMULA_SPILL_CELLS, +}; +use crate::SpreadsheetFormulaErrorLiteral; + +impl EvaluationContext<'_> { + pub(super) fn evaluate_structured_reference( + &self, + qualifier: Option<&SpreadsheetFormulaQualifier>, + reference: &str, + current: FormulaCellKey, + ) -> UseResult { + match self.tables.resolve( + qualifier, + reference, + current.sheet, + current.column, + current.row, + ) { + Ok(areas) => { + ensure_reference_area_count(areas.len())?; + Ok(EvalValue::Reference( + areas + .into_iter() + .map(|area| FormulaReferenceArea { + sheet: area.sheet, + start_column: area.start_column, + start_row: area.start_row, + end_column: area.end_column, + end_row: area.end_row, + }) + .collect(), + )) + } + Err(error) if matches!(error.kind, StructuredReferenceErrorKind::ExternalWorkbook) => { + Err(calculation_error( + "use.office.spreadsheet_formula_external_reference_unsupported", + error.message, + ) + .with_detail("reference", reference)) + } + Err(error) => Err(calculation_error( + "use.office.spreadsheet_formula_structured_reference_unsupported", + error.message, + ) + .with_detail("reference", reference)), + } + } + + pub(super) fn evaluate_reference( + &mut self, + reference: &SpreadsheetFormulaReference, + current: FormulaCellKey, + ) -> UseResult { + let Some(sheets) = self.resolve_reference_sheets(reference.qualifier.as_ref(), current)? + else { + return Ok(EvalValue::Scalar(ScalarValue::Error( + SpreadsheetFormulaErrorLiteral::Reference, + ))); + }; + let (start_column, start_row, end_column, end_row) = match reference.kind { + SpreadsheetFormulaReferenceKind::Cell { column, row, .. } => (column, row, column, row), + SpreadsheetFormulaReferenceKind::Column { column, .. } => { + (column, 1, column, crate::spreadsheet_reference::MAX_ROWS) + } + SpreadsheetFormulaReferenceKind::Row { row, .. } => { + (1, row, crate::spreadsheet_reference::MAX_COLUMNS, row) + } + }; + let areas = sheets + .into_iter() + .map(|sheet| FormulaReferenceArea { + sheet, + start_column, + start_row, + end_column, + end_row, + }) + .collect::>(); + ensure_reference_area_count(areas.len())?; + Ok(EvalValue::Reference(areas)) + } + + pub(super) fn evaluate_reference_operator( + &mut self, + operator: SpreadsheetFormulaBinaryOperator, + left: &SpreadsheetFormulaExpression, + right: &SpreadsheetFormulaExpression, + current: FormulaCellKey, + ) -> UseResult { + if matches!(operator, SpreadsheetFormulaBinaryOperator::Range) { + return self.evaluate_range(left, right, current); + } + let left = self.evaluate_expression(left, current)?; + let right = self.evaluate_expression(right, current)?; + let (EvalValue::Reference(mut left), EvalValue::Reference(right)) = (left, right) else { + return Ok(EvalValue::Scalar(ScalarValue::Error( + SpreadsheetFormulaErrorLiteral::Value, + ))); + }; + if matches!(operator, SpreadsheetFormulaBinaryOperator::Union) { + ensure_reference_area_total(left.len(), right.len())?; + left.extend(right); + return Ok(EvalValue::Reference(left)); + } + ensure_reference_comparisons(left.len(), right.len())?; + let mut areas = Vec::new(); + for left in left { + for right in &right { + if let Some(area) = left.intersect(*right) { + push_reference_area(&mut areas, area)?; + } + } + } + Ok(EvalValue::Reference(areas)) + } + + pub(super) fn evaluate_name( + &mut self, + qualifier: Option<&SpreadsheetFormulaQualifier>, + name: &str, + current: FormulaCellKey, + ) -> UseResult { + if qualifier.is_some_and(SpreadsheetFormulaQualifier::is_external) { + return Err(calculation_error( + "use.office.spreadsheet_formula_external_reference_unsupported", + format!( + "Native calculation never opens external reference '{}'.", + qualifier.map_or_else(|| name.to_string(), qualifier_label) + ), + )); + } + let explicit_scope = if let Some(qualifier) = qualifier { + if qualifier.is_three_dimensional() { + return Err(calculation_error( + "use.office.spreadsheet_formula_named_reference_unsupported", + format!( + "Native calculation does not resolve a name through 3D qualifier '{}'.", + qualifier_label(qualifier) + ), + )); + } + let Some(sheet) = self.sheet_position(&qualifier.worksheet) else { + return Ok(EvalValue::Scalar(ScalarValue::Error( + SpreadsheetFormulaErrorLiteral::Reference, + ))); + }; + Some(sheet) + } else { + None + }; + let normalized = name.to_lowercase(); + let local_scope = explicit_scope.or(Some(current.sheet)); + let definition = local_scope + .and_then(|sheet| { + self.named_definitions + .get(&(Some(sheet), normalized.clone())) + }) + .or_else(|| self.named_definitions.get(&(None, normalized.clone()))) + .cloned(); + let Some(definition) = definition else { + return Ok(EvalValue::Scalar(ScalarValue::Error( + SpreadsheetFormulaErrorLiteral::Name, + ))); + }; + if self.named_stack.len() >= MAX_SPREADSHEET_FORMULA_DEPTH { + return Err(calculation_error( + "use.office.spreadsheet_formula_named_reference_depth", + format!( + "Native calculation resolves at most {MAX_SPREADSHEET_FORMULA_DEPTH} nested named references." + ), + ) + .with_detail("namedRange", definition.name)); + } + let key = (definition.scope_sheet, normalized); + if !self.named_stack.insert(key.clone()) { + return Ok(EvalValue::Scalar(ScalarValue::Error( + SpreadsheetFormulaErrorLiteral::Name, + ))); + } + let parsed = parse_spreadsheet_formula(&definition.formula).map_err(|error| { + error + .with_detail("namedRange", definition.name.clone()) + .with_detail( + "scope", + definition.scope_sheet.map_or_else( + || "workbook".to_string(), + |sheet| { + self.sheet_names + .get(sheet) + .cloned() + .unwrap_or_else(|| "unknown".to_string()) + }, + ), + ) + })?; + let named_current = FormulaCellKey { + sheet: definition.scope_sheet.unwrap_or(current.sheet), + ..current + }; + let result = self.evaluate_expression(&parsed.root, named_current); + self.named_stack.remove(&key); + result + } + + pub(super) fn spill_reference(&self, value: EvalValue) -> UseResult { + let EvalValue::Reference(areas) = value else { + return Ok(EvalValue::Scalar(ScalarValue::Error( + SpreadsheetFormulaErrorLiteral::Reference, + ))); + }; + let Some(area) = areas.first().copied().filter(|_| areas.len() == 1) else { + return Ok(EvalValue::Scalar(ScalarValue::Error( + SpreadsheetFormulaErrorLiteral::Reference, + ))); + }; + if area.start_column != area.end_column || area.start_row != area.end_row { + return Ok(EvalValue::Scalar(ScalarValue::Error( + SpreadsheetFormulaErrorLiteral::Reference, + ))); + } + let anchor = FormulaCellKey { + sheet: area.sheet, + column: area.start_column, + row: area.start_row, + }; + Ok(self.spills.get(&anchor).copied().map_or_else( + || { + EvalValue::Scalar(ScalarValue::Error( + SpreadsheetFormulaErrorLiteral::Reference, + )) + }, + |spill| EvalValue::Reference(vec![spill]), + )) + } + + pub(super) fn implicit_intersection( + &self, + value: EvalValue, + current: FormulaCellKey, + ) -> UseResult { + match value { + EvalValue::Reference(areas) => { + let Some(area) = areas.first().copied() else { + return Ok(EvalValue::Scalar(ScalarValue::Error( + SpreadsheetFormulaErrorLiteral::Null, + ))); + }; + let key = if area.contains(current) { + current + } else if area.start_column == area.end_column + && (area.start_row..=area.end_row).contains(¤t.row) + { + FormulaCellKey { + sheet: area.sheet, + column: area.start_column, + row: current.row, + } + } else if area.start_row == area.end_row + && (area.start_column..=area.end_column).contains(¤t.column) + { + FormulaCellKey { + sheet: area.sheet, + column: current.column, + row: area.start_row, + } + } else { + FormulaCellKey { + sheet: area.sheet, + column: area.start_column, + row: area.start_row, + } + }; + Ok(EvalValue::Scalar( + self.values.get(&key).cloned().unwrap_or(ScalarValue::Blank), + )) + } + EvalValue::Array(array) => Ok(EvalValue::Scalar( + array + .rows + .first() + .and_then(|row| row.first()) + .cloned() + .unwrap_or(ScalarValue::Blank), + )), + value => Ok(value), + } + } + + pub(super) fn materialize_areas(&self, areas: &[FormulaReferenceArea]) -> UseResult { + ensure_reference_area_count(areas.len())?; + if areas.is_empty() { + return Ok(EvalValue::Scalar(ScalarValue::Error( + SpreadsheetFormulaErrorLiteral::Null, + ))); + } + if areas.len() > 1 { + let mut rows = Vec::new(); + let mut text_bytes = 0_usize; + for area in areas { + let cells = area.cell_count().ok_or_else(reference_limit_error)?; + if cells > MAX_SPREADSHEET_FORMULA_SPILL_CELLS + || rows.len().saturating_add(cells) > MAX_SPREADSHEET_FORMULA_SPILL_CELLS + { + return Err(reference_limit_error()); + } + for row in area.start_row..=area.end_row { + for column in area.start_column..=area.end_column { + let value = self + .values + .get(&FormulaCellKey { + sheet: area.sheet, + column, + row, + }) + .cloned() + .unwrap_or(ScalarValue::Blank); + accumulate_formula_text_bytes(&mut text_bytes, &value)?; + rows.push(vec![value]); + } + } + } + return Ok(EvalValue::Array( + FormulaArray::new(rows).ok_or_else(reference_limit_error)?, + )); + } + let area = areas.first().copied().ok_or_else(invalid_array_shape)?; + let cells = area.cell_count().ok_or_else(reference_limit_error)?; + if cells > MAX_SPREADSHEET_FORMULA_SPILL_CELLS { + return Err(reference_limit_error().with_detail("cells", cells)); + } + if cells == 1 { + return Ok(EvalValue::Scalar( + self.values + .get(&FormulaCellKey { + sheet: area.sheet, + column: area.start_column, + row: area.start_row, + }) + .cloned() + .unwrap_or(ScalarValue::Blank), + )); + } + let mut rows = Vec::new(); + let mut text_bytes = 0_usize; + for row in area.start_row..=area.end_row { + let mut values = Vec::new(); + for column in area.start_column..=area.end_column { + let value = self + .values + .get(&FormulaCellKey { + sheet: area.sheet, + column, + row, + }) + .cloned() + .unwrap_or(ScalarValue::Blank); + accumulate_formula_text_bytes(&mut text_bytes, &value)?; + values.push(value); + } + rows.push(values); + } + Ok(EvalValue::Array(FormulaArray { rows })) + } + + pub(super) fn populated_reference_values( + &self, + areas: &[FormulaReferenceArea], + ) -> UseResult> { + ensure_reference_area_count(areas.len())?; + let mut values = Vec::new(); + for area in areas { + let mut area_values = Vec::new(); + for (key, value) in self.values.iter().filter(|(key, _)| area.contains(**key)) { + let count = values + .len() + .checked_add(area_values.len()) + .and_then(|count| count.checked_add(1)) + .ok_or_else(reference_limit_error)?; + if count > MAX_SPREADSHEET_FORMULA_SPILL_CELLS { + return Err(reference_limit_error().with_detail("cells", count)); + } + area_values.push((*key, value.clone())); + } + area_values.sort_by_key(|(key, _)| (key.row, key.column)); + values.extend(area_values.into_iter().map(|(_, value)| value)); + } + Ok(values) + } + + fn evaluate_range( + &mut self, + left: &SpreadsheetFormulaExpression, + right: &SpreadsheetFormulaExpression, + current: FormulaCellKey, + ) -> UseResult { + let (Some(left), Some(right)) = (endpoint_reference(left), endpoint_reference(right)) + else { + return Ok(EvalValue::Scalar(ScalarValue::Error( + SpreadsheetFormulaErrorLiteral::Reference, + ))); + }; + let qualifier = match (&left.qualifier, &right.qualifier) { + (Some(left), Some(right)) if left != right => { + return Ok(EvalValue::Scalar(ScalarValue::Error( + SpreadsheetFormulaErrorLiteral::Reference, + ))); + } + (Some(qualifier), _) | (_, Some(qualifier)) => Some(qualifier), + (None, None) => None, + }; + let Some(sheets) = self.resolve_reference_sheets(qualifier, current)? else { + return Ok(EvalValue::Scalar(ScalarValue::Error( + SpreadsheetFormulaErrorLiteral::Reference, + ))); + }; + let coordinates = match (left.kind, right.kind) { + ( + SpreadsheetFormulaReferenceKind::Cell { + column: left_column, + row: left_row, + .. + }, + SpreadsheetFormulaReferenceKind::Cell { + column: right_column, + row: right_row, + .. + }, + ) => ( + left_column.min(right_column), + left_row.min(right_row), + left_column.max(right_column), + left_row.max(right_row), + ), + ( + SpreadsheetFormulaReferenceKind::Column { + column: left_column, + .. + }, + SpreadsheetFormulaReferenceKind::Column { + column: right_column, + .. + }, + ) => ( + left_column.min(right_column), + 1, + left_column.max(right_column), + crate::spreadsheet_reference::MAX_ROWS, + ), + ( + SpreadsheetFormulaReferenceKind::Row { row: left_row, .. }, + SpreadsheetFormulaReferenceKind::Row { row: right_row, .. }, + ) => ( + 1, + left_row.min(right_row), + crate::spreadsheet_reference::MAX_COLUMNS, + left_row.max(right_row), + ), + _ => { + return Ok(EvalValue::Scalar(ScalarValue::Error( + SpreadsheetFormulaErrorLiteral::Reference, + ))); + } + }; + let areas = sheets + .into_iter() + .map(|sheet| FormulaReferenceArea { + sheet, + start_column: coordinates.0, + start_row: coordinates.1, + end_column: coordinates.2, + end_row: coordinates.3, + }) + .collect::>(); + ensure_reference_area_count(areas.len())?; + Ok(EvalValue::Reference(areas)) + } + + fn resolve_reference_sheets( + &self, + qualifier: Option<&SpreadsheetFormulaQualifier>, + current: FormulaCellKey, + ) -> UseResult>> { + let Some(qualifier) = qualifier else { + return Ok(Some(vec![current.sheet])); + }; + if qualifier.is_external() { + return Err(calculation_error( + "use.office.spreadsheet_formula_external_reference_unsupported", + format!( + "Native calculation never opens external reference '{}'.", + qualifier_label(qualifier) + ), + )); + } + let Some(start) = self.sheet_position(&qualifier.worksheet) else { + return Ok(None); + }; + let Some(end_name) = qualifier.worksheet_end.as_deref() else { + return Ok(Some(vec![start])); + }; + let Some(end) = self.sheet_position(end_name) else { + return Ok(None); + }; + let low = start.min(end); + let high = start.max(end); + let areas = high + .checked_sub(low) + .and_then(|distance| distance.checked_add(1)) + .ok_or_else(reference_area_limit_error)?; + ensure_reference_area_count(areas)?; + Ok(Some((low..=high).collect())) + } + + fn sheet_position(&self, name: &str) -> Option { + self.sheet_names + .iter() + .position(|sheet| sheet.eq_ignore_ascii_case(name)) + } +} + +fn endpoint_reference( + expression: &SpreadsheetFormulaExpression, +) -> Option<&SpreadsheetFormulaReference> { + match &expression.kind { + SpreadsheetFormulaExpressionKind::Reference(reference) => Some(reference), + SpreadsheetFormulaExpressionKind::Parenthesized(inner) => endpoint_reference(inner), + _ => None, + } +} + +fn qualifier_label(qualifier: &SpreadsheetFormulaQualifier) -> String { + let workbook = qualifier.workbook.as_deref().unwrap_or_default(); + let worksheet = qualifier.worksheet_end.as_ref().map_or_else( + || qualifier.worksheet.clone(), + |end| format!("{}:{end}", qualifier.worksheet), + ); + format!("{workbook}{worksheet}") +} + +fn reference_limit_error() -> a3s_use_core::UseError { + calculation_error( + "use.office.spreadsheet_formula_reference_limit", + format!( + "Native calculation materializes at most {MAX_SPREADSHEET_FORMULA_SPILL_CELLS} referenced cells per value." + ), + ) +} + +fn push_reference_area( + areas: &mut Vec, + area: FormulaReferenceArea, +) -> UseResult<()> { + let total = areas + .len() + .checked_add(1) + .ok_or_else(reference_area_limit_error)?; + ensure_reference_area_count(total)?; + areas.push(area); + Ok(()) +} + +fn ensure_reference_area_total(left: usize, right: usize) -> UseResult<()> { + let total = left + .checked_add(right) + .ok_or_else(reference_area_limit_error)?; + ensure_reference_area_count(total) +} + +fn ensure_reference_area_count(areas: usize) -> UseResult<()> { + if areas > MAX_SPREADSHEET_FORMULA_REFERENCE_AREAS { + return Err(reference_area_limit_error().with_detail("areas", areas)); + } + Ok(()) +} + +fn ensure_reference_comparisons(left: usize, right: usize) -> UseResult<()> { + let visits = left + .checked_mul(right) + .ok_or_else(reference_visit_limit_error)?; + if visits > MAX_SPREADSHEET_FORMULA_REFERENCE_VISITS { + return Err(reference_visit_limit_error().with_detail("visits", visits)); + } + Ok(()) +} + +fn reference_area_limit_error() -> a3s_use_core::UseError { + calculation_error( + "use.office.spreadsheet_formula_reference_area_limit", + format!( + "Native calculation retains at most {MAX_SPREADSHEET_FORMULA_REFERENCE_AREAS} reference areas per value." + ), + ) +} + +fn reference_visit_limit_error() -> a3s_use_core::UseError { + calculation_error( + "use.office.spreadsheet_formula_reference_visit_limit", + format!( + "Native calculation reference operators visit at most {MAX_SPREADSHEET_FORMULA_REFERENCE_VISITS} area pairs." + ), + ) +} diff --git a/crates/office/src/spreadsheet_formula/evaluate/spill.rs b/crates/office/src/spreadsheet_formula/evaluate/spill.rs new file mode 100644 index 00000000..32c85f1b --- /dev/null +++ b/crates/office/src/spreadsheet_formula/evaluate/spill.rs @@ -0,0 +1,191 @@ +use a3s_use_core::UseResult; + +use crate::spreadsheet_reference::{CellReference, MAX_COLUMNS, MAX_ROWS}; +use crate::{SpreadsheetFormulaErrorLiteral, SpreadsheetFormulaValue}; + +use super::{ + calculation_error, calculation_text_limit_error, ensure_formula_text_limit, public_scalar, + spill_limit_error, EvalValue, EvaluationContext, FormulaCellKey, FormulaReferenceArea, + ScalarValue, MAX_SPREADSHEET_FORMULA_CALCULATION_TEXT_BYTES, + MAX_SPREADSHEET_FORMULA_SPILL_CELLS, +}; + +impl EvaluationContext<'_> { + pub(super) fn finalize_formula_result( + &mut self, + anchor: FormulaCellKey, + value: EvalValue, + existing_spill_cells: usize, + existing_text_bytes: usize, + ) -> UseResult<(SpreadsheetFormulaValue, Option, usize, usize)> { + self.clear_previous_spill(anchor); + let materialized = self.materialize(value)?; + match materialized { + EvalValue::Scalar(value) => { + ensure_formula_text_limit(&value)?; + let text_bytes = scalar_text_bytes(&value); + ensure_calculation_text_limit(existing_text_bytes, text_bytes)?; + self.values.insert(anchor, value.clone()); + Ok((public_scalar(&value), None, 0, text_bytes)) + } + EvalValue::Array(array) => { + let mut text_bytes = 0_usize; + for value in array.rows.iter().flatten() { + ensure_formula_text_limit(value)?; + text_bytes = text_bytes + .checked_add(scalar_text_bytes(value)) + .ok_or_else(calculation_text_limit_error)?; + } + let height = u32::try_from(array.height()).map_err(|_| spill_limit_error())?; + let width = u32::try_from(array.width()).map_err(|_| spill_limit_error())?; + let cells = array + .height() + .checked_mul(array.width()) + .ok_or_else(spill_limit_error)?; + if cells > MAX_SPREADSHEET_FORMULA_SPILL_CELLS { + return Err(spill_limit_error().with_detail("cells", cells)); + } + let spill_cells = cells.saturating_sub(1); + let end_row = anchor + .row + .checked_add(height.saturating_sub(1)) + .filter(|row| *row <= MAX_ROWS); + let end_column = anchor + .column + .checked_add(width.saturating_sub(1)) + .filter(|column| *column <= MAX_COLUMNS); + let (Some(end_row), Some(end_column)) = (end_row, end_column) else { + let error = ScalarValue::Error(SpreadsheetFormulaErrorLiteral::Spill); + self.values.insert(anchor, error.clone()); + return Ok((public_scalar(&error), None, 0, 0)); + }; + let area = FormulaReferenceArea { + sheet: anchor.sheet, + start_column: anchor.column, + start_row: anchor.row, + end_column, + end_row, + }; + if self.spill_is_blocked(anchor, area) { + let error = ScalarValue::Error(SpreadsheetFormulaErrorLiteral::Spill); + self.values.insert(anchor, error.clone()); + return Ok((public_scalar(&error), None, 0, 0)); + } + let total_spill_cells = existing_spill_cells + .checked_add(spill_cells) + .ok_or_else(spill_limit_error)?; + if total_spill_cells > MAX_SPREADSHEET_FORMULA_SPILL_CELLS { + return Err(spill_limit_error().with_detail("cells", total_spill_cells)); + } + ensure_calculation_text_limit(existing_text_bytes, text_bytes)?; + for (row_offset, row) in array.rows.iter().enumerate() { + for (column_offset, value) in row.iter().enumerate() { + let key = FormulaCellKey { + sheet: anchor.sheet, + column: anchor.column + + u32::try_from(column_offset).map_err(|_| spill_limit_error())?, + row: anchor.row + + u32::try_from(row_offset).map_err(|_| spill_limit_error())?, + }; + self.values.insert(key, value.clone()); + if key != anchor { + self.spill_owners.insert(key, anchor); + } + } + } + self.spills.insert(anchor, area); + let start = CellReference { + column: anchor.column, + row: anchor.row, + } + .a1(); + let end = CellReference { + column: end_column, + row: end_row, + } + .a1(); + Ok(( + SpreadsheetFormulaValue::Array { + rows: array + .rows + .iter() + .map(|row| row.iter().map(public_scalar).collect()) + .collect(), + }, + Some(if start == end { + start + } else { + format!("{start}:{end}") + }), + spill_cells, + text_bytes, + )) + } + EvalValue::Reference(_) => Err(calculation_error( + "use.office.spreadsheet_formula_calculation_invalid", + "Formula result retained an unresolved reference.", + )), + } + } + + fn clear_previous_spill(&mut self, anchor: FormulaCellKey) { + let Some(area) = self.old_spills.get(&anchor).copied() else { + return; + }; + for row in area.start_row..=area.end_row { + for column in area.start_column..=area.end_column { + let key = FormulaCellKey { + sheet: area.sheet, + column, + row, + }; + if key != anchor { + self.values.remove(&key); + } + } + } + } + + fn spill_is_blocked(&self, anchor: FormulaCellKey, area: FormulaReferenceArea) -> bool { + for row in area.start_row..=area.end_row { + for column in area.start_column..=area.end_column { + let key = FormulaCellKey { + sheet: area.sheet, + column, + row, + }; + if key == anchor { + continue; + } + if self.formula_cells.contains(&key) || self.spill_owners.contains_key(&key) { + return true; + } + let owned_before = self + .old_spills + .get(&anchor) + .is_some_and(|old| old.contains(key)); + if self.occupied.contains(&key) && !owned_before { + return true; + } + } + } + false + } +} + +fn scalar_text_bytes(value: &ScalarValue) -> usize { + match value { + ScalarValue::Text(value) => value.len(), + _ => 0, + } +} + +fn ensure_calculation_text_limit(existing: usize, added: usize) -> UseResult<()> { + let bytes = existing + .checked_add(added) + .ok_or_else(calculation_text_limit_error)?; + if bytes > MAX_SPREADSHEET_FORMULA_CALCULATION_TEXT_BYTES { + return Err(calculation_text_limit_error().with_detail("bytes", bytes)); + } + Ok(()) +} diff --git a/crates/office/src/spreadsheet_formula/graph.rs b/crates/office/src/spreadsheet_formula/graph.rs new file mode 100644 index 00000000..6a9ab4df --- /dev/null +++ b/crates/office/src/spreadsheet_formula/graph.rs @@ -0,0 +1,515 @@ +mod reference; + +use std::collections::{BTreeMap, BTreeSet}; + +use a3s_use_core::{UseError, UseResult}; +use serde::{Deserialize, Serialize}; + +use crate::discovery::office_error; +use crate::semantic::{DocumentNode, NativeOfficeDocument, OfficeNodeType}; +use crate::spreadsheet_reference::CellReference; +use crate::DocumentKind; + +use super::{parse_spreadsheet_formula, SpreadsheetFormula}; +use crate::spreadsheet_formula::structured_reference::FormulaTableCatalog; +use reference::{ + collect_references, FormulaNamedDefinition, FormulaReferenceArea, FormulaReferenceCollection, +}; + +/// Maximum formula cells admitted to one native dependency graph. +pub const MAX_SPREADSHEET_FORMULA_CELLS: usize = 100_000; + +/// Maximum formula-to-formula edges admitted to one dependency graph. +pub const MAX_SPREADSHEET_FORMULA_DEPENDENCIES: usize = 1_000_000; + +/// Maximum formula-cell candidates visited while resolving all static +/// reference areas in one dependency graph. +pub const MAX_SPREADSHEET_FORMULA_REFERENCE_VISITS: usize = 1_000_000; + +/// Stable workbook coordinate for a formula cell. +#[derive(Debug, Clone, PartialEq, Eq, PartialOrd, Ord, Hash, Serialize, Deserialize)] +#[serde(rename_all = "camelCase")] +pub struct SpreadsheetFormulaCell { + pub sheet: String, + pub column: u32, + pub row: u32, +} + +impl SpreadsheetFormulaCell { + pub fn path(&self) -> String { + format!( + "/{}/{}{}", + self.sheet, + crate::spreadsheet_reference::column_name(self.column), + self.row + ) + } +} + +/// Reason a static dependency could not be resolved to workbook cells. +#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Serialize, Deserialize)] +#[serde(rename_all = "kebab-case")] +pub enum SpreadsheetFormulaUnresolvedReferenceKind { + MissingWorksheet, + ExternalWorkbook, + UndefinedName, + NamedRangeCycle, + NamedRangeDepth, + StructuredReference, + DynamicReference, + UnsupportedReference, +} + +/// One source reference that cannot participate in the static graph. +#[derive(Debug, Clone, PartialEq, Eq, PartialOrd, Ord, Serialize, Deserialize)] +#[serde(rename_all = "camelCase")] +pub struct SpreadsheetFormulaUnresolvedReference { + pub kind: SpreadsheetFormulaUnresolvedReferenceKind, + pub reference: String, +} + +/// One formula cell and its formula-to-formula graph edges. +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +#[serde(rename_all = "camelCase")] +pub struct SpreadsheetFormulaDependencyNode { + pub cell: SpreadsheetFormulaCell, + pub formula: String, + pub dependencies: Vec, + pub dependents: Vec, + #[serde(default, skip_serializing_if = "Vec::is_empty")] + pub unresolved_references: Vec, +} + +/// Bounded static dependency graph and deterministic calculation order. +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +#[serde(rename_all = "camelCase")] +pub struct SpreadsheetFormulaDependencyGraph { + pub nodes: Vec, + pub calculation_order: Vec, + #[serde(default, skip_serializing_if = "Vec::is_empty")] + pub cycles: Vec>, +} + +impl SpreadsheetFormulaDependencyGraph { + pub fn is_acyclic(&self) -> bool { + self.cycles.is_empty() + } +} + +impl NativeOfficeDocument { + /// Builds a bounded formula-to-formula dependency graph without evaluating + /// any formula or fetching an external workbook. + pub fn formula_dependency_graph(&self) -> UseResult { + build_dependency_graph(self) + } +} + +struct FormulaRecord { + cell: SpreadsheetFormulaCell, + formula: String, + parsed: SpreadsheetFormula, + sheet_index: usize, +} + +fn build_dependency_graph( + document: &NativeOfficeDocument, +) -> UseResult { + if document.kind() != DocumentKind::Spreadsheet { + return Err(graph_error( + "use.office.spreadsheet_formula_graph_type_unsupported", + "Formula dependency graphs are available only for Spreadsheet documents.", + )); + } + let sheet_names = document + .root() + .children + .iter() + .filter(|node| node.node_type == OfficeNodeType::Worksheet) + .map(|node| node.path.trim_start_matches('/').to_string()) + .collect::>(); + let tables = FormulaTableCatalog::collect(document.root(), &sheet_names)?; + let named_definitions = collect_named_definitions(document.root(), &sheet_names); + let records = collect_formula_records(document.root(), &sheet_names)?; + let mut rows_by_sheet = vec![BTreeMap::>::new(); sheet_names.len()]; + for (index, record) in records.iter().enumerate() { + rows_by_sheet + .get_mut(record.sheet_index) + .ok_or_else(graph_index_error)? + .entry(record.cell.row) + .or_default() + .insert(record.cell.column, index); + } + + let mut dependencies = vec![BTreeSet::::new(); records.len()]; + let mut unresolved = Vec::with_capacity(records.len()); + let mut edge_count = 0_usize; + let mut reference_visits = 0_usize; + for (index, record) in records.iter().enumerate() { + let FormulaReferenceCollection { + areas, + mut unresolved_references, + } = collect_references( + &record.parsed.root, + record.sheet_index, + record.cell.column, + record.cell.row, + &sheet_names, + &named_definitions, + &tables, + )?; + unresolved_references.sort(); + unresolved_references.dedup(); + unresolved.push(unresolved_references); + for area in areas { + for dependency in formula_cells_in_area(&area, &rows_by_sheet, &mut reference_visits)? { + if dependencies + .get_mut(index) + .ok_or_else(graph_index_error)? + .insert(dependency) + { + edge_count = edge_count.saturating_add(1); + if edge_count > MAX_SPREADSHEET_FORMULA_DEPENDENCIES { + return Err(graph_error( + "use.office.spreadsheet_formula_dependency_limit", + format!( + "Spreadsheet formula graph exceeds {MAX_SPREADSHEET_FORMULA_DEPENDENCIES} formula dependencies." + ), + ) + .with_detail("dependencies", edge_count)); + } + } + } + } + } + + let mut dependents = vec![BTreeSet::::new(); records.len()]; + for (cell, cell_dependencies) in dependencies.iter().enumerate() { + for dependency in cell_dependencies { + dependents + .get_mut(*dependency) + .ok_or_else(graph_index_error)? + .insert(cell); + } + } + let (calculation_order, remaining) = topological_order(&dependencies, &dependents)?; + let cycles = strongly_connected_cycles(&dependencies, &dependents, &remaining)?; + + let mut nodes = Vec::with_capacity(records.len()); + for (index, record) in records.iter().enumerate() { + let cell_dependencies = dependencies.get(index).ok_or_else(graph_index_error)?; + let cell_dependents = dependents.get(index).ok_or_else(graph_index_error)?; + let unresolved_references = unresolved.get(index).ok_or_else(graph_index_error)?; + nodes.push(SpreadsheetFormulaDependencyNode { + cell: record.cell.clone(), + formula: record.formula.clone(), + dependencies: cell_dependencies + .iter() + .map(|dependency| { + records + .get(*dependency) + .map(|record| record.cell.clone()) + .ok_or_else(graph_index_error) + }) + .collect::>>()?, + dependents: cell_dependents + .iter() + .map(|dependent| { + records + .get(*dependent) + .map(|record| record.cell.clone()) + .ok_or_else(graph_index_error) + }) + .collect::>>()?, + unresolved_references: unresolved_references.clone(), + }); + } + Ok(SpreadsheetFormulaDependencyGraph { + nodes, + calculation_order: calculation_order + .into_iter() + .map(|index| { + records + .get(index) + .map(|record| record.cell.clone()) + .ok_or_else(graph_index_error) + }) + .collect::>>()?, + cycles: cycles + .into_iter() + .map(|cycle| { + cycle + .into_iter() + .map(|index| { + records + .get(index) + .map(|record| record.cell.clone()) + .ok_or_else(graph_index_error) + }) + .collect::>>() + }) + .collect::>>()?, + }) +} + +fn collect_formula_records( + root: &DocumentNode, + sheet_names: &[String], +) -> UseResult> { + let mut records = Vec::new(); + for (sheet_index, sheet_name) in sheet_names.iter().enumerate() { + let sheet = root + .children + .iter() + .find(|node| { + node.node_type == OfficeNodeType::Worksheet + && node + .path + .strip_prefix('/') + .is_some_and(|path| path.eq_ignore_ascii_case(sheet_name)) + }) + .ok_or_else(|| { + graph_error( + "use.office.spreadsheet_formula_graph_invalid", + format!("Formula graph cannot find worksheet '{sheet_name}'."), + ) + })?; + for row in &sheet.children { + if row.node_type != OfficeNodeType::Row { + continue; + } + for cell in &row.children { + if cell.node_type != OfficeNodeType::Cell { + continue; + } + let Some(formula) = cell.format.get("formula") else { + continue; + }; + if records.len() >= MAX_SPREADSHEET_FORMULA_CELLS { + return Err(graph_error( + "use.office.spreadsheet_formula_cell_limit", + format!( + "Spreadsheet formula graph accepts at most {MAX_SPREADSHEET_FORMULA_CELLS} formula cells." + ), + ) + .with_detail("formulas", records.len().saturating_add(1))); + } + let reference = cell + .path + .rsplit_once('/') + .map(|(_, reference)| reference) + .ok_or_else(|| invalid_formula_cell(cell))?; + let reference = + CellReference::parse(reference).map_err(|_| invalid_formula_cell(cell))?; + let parsed = parse_spreadsheet_formula(formula) + .map_err(|error| formula_parse_error(error, &cell.path))?; + records.push(FormulaRecord { + cell: SpreadsheetFormulaCell { + sheet: sheet_name.clone(), + column: reference.column, + row: reference.row, + }, + formula: formula.clone(), + parsed, + sheet_index, + }); + } + } + } + Ok(records) +} + +fn collect_named_definitions( + root: &DocumentNode, + sheet_names: &[String], +) -> BTreeMap<(Option, String), FormulaNamedDefinition> { + let mut definitions = BTreeMap::new(); + for node in root + .children + .iter() + .filter(|node| node.node_type == OfficeNodeType::NamedRangeCollection) + .flat_map(|collection| &collection.children) + .filter(|node| node.node_type == OfficeNodeType::NamedRange) + { + let (Some(name), Some(reference), Some(scope)) = ( + node.format.get("name"), + node.format.get("ref"), + node.format.get("scope"), + ) else { + continue; + }; + let scope_sheet = if scope.eq_ignore_ascii_case("workbook") { + None + } else { + let sheet = scope.strip_prefix("worksheet:").unwrap_or(scope); + sheet_names + .iter() + .position(|candidate| candidate.eq_ignore_ascii_case(sheet)) + }; + definitions.insert( + (scope_sheet, name.to_lowercase()), + FormulaNamedDefinition { + name: name.clone(), + formula: reference.clone(), + scope_sheet, + }, + ); + } + definitions +} + +fn formula_cells_in_area( + area: &FormulaReferenceArea, + rows_by_sheet: &[BTreeMap>], + reference_visits: &mut usize, +) -> UseResult> { + let Some(rows) = rows_by_sheet.get(area.sheet_index) else { + return Ok(Vec::new()); + }; + let mut cells = Vec::new(); + for (_, columns) in rows.range(area.start_row..=area.end_row) { + for (_, index) in columns.range(area.start_column..=area.end_column) { + *reference_visits = reference_visits + .checked_add(1) + .ok_or_else(reference_visit_limit)?; + if *reference_visits > MAX_SPREADSHEET_FORMULA_REFERENCE_VISITS { + return Err(reference_visit_limit().with_detail("visits", *reference_visits)); + } + cells.push(*index); + } + } + Ok(cells) +} + +fn topological_order( + dependencies: &[BTreeSet], + dependents: &[BTreeSet], +) -> UseResult<(Vec, BTreeSet)> { + let mut indegree = dependencies.iter().map(BTreeSet::len).collect::>(); + let mut ready = indegree + .iter() + .enumerate() + .filter_map(|(index, count)| (*count == 0).then_some(index)) + .collect::>(); + let mut order = Vec::with_capacity(dependencies.len()); + while let Some(index) = ready.first().copied() { + ready.remove(&index); + order.push(index); + for dependent in dependents.get(index).ok_or_else(graph_index_error)? { + let degree = indegree.get_mut(*dependent).ok_or_else(graph_index_error)?; + *degree = degree.saturating_sub(1); + if *degree == 0 { + ready.insert(*dependent); + } + } + } + let remaining = (0..dependencies.len()) + .filter(|index| indegree.get(*index).is_some_and(|degree| *degree > 0)) + .collect(); + Ok((order, remaining)) +} + +fn strongly_connected_cycles( + dependencies: &[BTreeSet], + dependents: &[BTreeSet], + remaining: &BTreeSet, +) -> UseResult>> { + let mut visited = BTreeSet::new(); + let mut finish = Vec::new(); + for start in remaining { + if visited.contains(start) { + continue; + } + let mut stack = vec![(*start, false)]; + while let Some((node, expanded)) = stack.pop() { + if expanded { + finish.push(node); + continue; + } + if !visited.insert(node) { + continue; + } + stack.push((node, true)); + for dependency in dependencies + .get(node) + .ok_or_else(graph_index_error)? + .iter() + .rev() + { + if remaining.contains(dependency) && !visited.contains(dependency) { + stack.push((*dependency, false)); + } + } + } + } + + visited.clear(); + let mut cycles = Vec::new(); + for start in finish.into_iter().rev() { + if visited.contains(&start) { + continue; + } + let mut component = Vec::new(); + let mut stack = vec![start]; + visited.insert(start); + while let Some(node) = stack.pop() { + component.push(node); + for dependent in dependents + .get(node) + .ok_or_else(graph_index_error)? + .iter() + .rev() + { + if remaining.contains(dependent) && visited.insert(*dependent) { + stack.push(*dependent); + } + } + } + component.sort_unstable(); + let first = component.first().copied().ok_or_else(graph_index_error)?; + if component.len() > 1 + || dependencies + .get(first) + .ok_or_else(graph_index_error)? + .contains(&first) + { + cycles.push(component); + } + } + cycles.sort_by_key(|cycle| cycle.first().copied().unwrap_or(usize::MAX)); + Ok(cycles) +} + +fn invalid_formula_cell(cell: &DocumentNode) -> UseError { + graph_error( + "use.office.spreadsheet_formula_graph_invalid", + format!( + "Formula cell '{}' has an invalid semantic coordinate.", + cell.path + ), + ) + .with_detail("cell", cell.path.clone()) +} + +fn formula_parse_error(error: UseError, cell: &str) -> UseError { + error.with_detail("cell", cell.to_string()) +} + +fn graph_error(code: &str, message: impl Into) -> UseError { + office_error(code, message) +} + +fn graph_index_error() -> UseError { + graph_error( + "use.office.spreadsheet_formula_graph_invalid", + "Spreadsheet formula graph contains an inconsistent internal index.", + ) +} + +fn reference_visit_limit() -> UseError { + graph_error( + "use.office.spreadsheet_formula_reference_visit_limit", + format!( + "Spreadsheet formula graph visits at most {MAX_SPREADSHEET_FORMULA_REFERENCE_VISITS} formula-cell reference candidates." + ), + ) +} diff --git a/crates/office/src/spreadsheet_formula/graph/reference.rs b/crates/office/src/spreadsheet_formula/graph/reference.rs new file mode 100644 index 00000000..d28d27bd --- /dev/null +++ b/crates/office/src/spreadsheet_formula/graph/reference.rs @@ -0,0 +1,582 @@ +use std::collections::{BTreeMap, BTreeSet}; + +use a3s_use_core::UseResult; + +use crate::spreadsheet_reference::{MAX_COLUMNS, MAX_ROWS}; + +use super::{ + graph_error, SpreadsheetFormulaUnresolvedReference, SpreadsheetFormulaUnresolvedReferenceKind, + MAX_SPREADSHEET_FORMULA_REFERENCE_VISITS, +}; +use crate::spreadsheet_formula::structured_reference::{ + FormulaTableCatalog, StructuredReferenceErrorKind, +}; +use crate::spreadsheet_formula::{ + parse_spreadsheet_formula, SpreadsheetFormulaBinaryOperator, SpreadsheetFormulaExpression, + SpreadsheetFormulaExpressionKind, SpreadsheetFormulaPostfixOperator, + SpreadsheetFormulaQualifier, SpreadsheetFormulaReference, SpreadsheetFormulaReferenceKind, + SpreadsheetFormulaUnaryOperator, MAX_SPREADSHEET_FORMULA_DEPTH, + MAX_SPREADSHEET_FORMULA_REFERENCE_AREAS, +}; + +#[derive(Debug, Clone)] +pub(super) struct FormulaNamedDefinition { + pub(super) name: String, + pub(super) formula: String, + pub(super) scope_sheet: Option, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord)] +pub(super) struct FormulaReferenceArea { + pub(super) sheet_index: usize, + pub(super) start_column: u32, + pub(super) start_row: u32, + pub(super) end_column: u32, + pub(super) end_row: u32, +} + +impl FormulaReferenceArea { + fn intersect(self, other: Self) -> Option { + if self.sheet_index != other.sheet_index { + return None; + } + let start_column = self.start_column.max(other.start_column); + let start_row = self.start_row.max(other.start_row); + let end_column = self.end_column.min(other.end_column); + let end_row = self.end_row.min(other.end_row); + (start_column <= end_column && start_row <= end_row).then_some(Self { + sheet_index: self.sheet_index, + start_column, + start_row, + end_column, + end_row, + }) + } +} + +pub(super) struct FormulaReferenceCollection { + pub(super) areas: Vec, + pub(super) unresolved_references: Vec, +} + +pub(super) fn collect_references( + expression: &SpreadsheetFormulaExpression, + current_sheet: usize, + current_column: u32, + current_row: u32, + sheet_names: &[String], + named_definitions: &BTreeMap<(Option, String), FormulaNamedDefinition>, + tables: &FormulaTableCatalog, +) -> UseResult { + let mut collector = ReferenceCollector { + sheet_names, + named_definitions, + tables, + current_column, + current_row, + areas: Vec::new(), + unresolved: Vec::new(), + visited_names: BTreeSet::new(), + }; + collector.collect_expression(expression, current_sheet)?; + collector.areas.sort(); + collector.areas.dedup(); + Ok(FormulaReferenceCollection { + areas: collector.areas, + unresolved_references: collector.unresolved, + }) +} + +struct ReferenceCollector<'a> { + sheet_names: &'a [String], + named_definitions: &'a BTreeMap<(Option, String), FormulaNamedDefinition>, + tables: &'a FormulaTableCatalog, + current_column: u32, + current_row: u32, + areas: Vec, + unresolved: Vec, + visited_names: BTreeSet<(Option, String)>, +} + +impl ReferenceCollector<'_> { + fn collect_expression( + &mut self, + expression: &SpreadsheetFormulaExpression, + current_sheet: usize, + ) -> UseResult<()> { + if let Some(areas) = self.try_reference_areas(expression, current_sheet)? { + self.extend_areas(areas)?; + return Ok(()); + } + match &expression.kind { + SpreadsheetFormulaExpressionKind::Name { qualifier, name } => { + self.collect_named_reference(qualifier.as_ref(), name, current_sheet) + } + SpreadsheetFormulaExpressionKind::StructuredReference { + qualifier, + reference, + } => { + match self.tables.resolve( + qualifier.as_ref(), + reference, + current_sheet, + self.current_column, + self.current_row, + ) { + Ok(areas) => self.extend_areas( + areas + .into_iter() + .map(|area| FormulaReferenceArea { + sheet_index: area.sheet, + start_column: area.start_column, + start_row: area.start_row, + end_column: area.end_column, + end_row: area.end_row, + }) + .collect(), + )?, + Err(error) => { + let kind = + if matches!(error.kind, StructuredReferenceErrorKind::ExternalWorkbook) + { + SpreadsheetFormulaUnresolvedReferenceKind::ExternalWorkbook + } else { + SpreadsheetFormulaUnresolvedReferenceKind::StructuredReference + }; + self.push_unresolved(kind, reference); + } + } + Ok(()) + } + SpreadsheetFormulaExpressionKind::Unary { operand, .. } + | SpreadsheetFormulaExpressionKind::Postfix { operand, .. } + | SpreadsheetFormulaExpressionKind::Parenthesized(operand) => { + self.collect_expression(operand, current_sheet) + } + SpreadsheetFormulaExpressionKind::Binary { + operator, + left, + right, + } => { + if matches!(operator, SpreadsheetFormulaBinaryOperator::Range) { + self.push_unresolved( + SpreadsheetFormulaUnresolvedReferenceKind::UnsupportedReference, + "range with non-A1 endpoints", + ); + } + self.collect_expression(left, current_sheet)?; + self.collect_expression(right, current_sheet) + } + SpreadsheetFormulaExpressionKind::FunctionCall { + name, arguments, .. + } => { + if matches!( + name.to_ascii_uppercase().as_str(), + "INDIRECT" | "OFFSET" | "INDEX" + ) { + self.push_unresolved( + SpreadsheetFormulaUnresolvedReferenceKind::DynamicReference, + name, + ); + } + for argument in arguments.iter().flatten() { + self.collect_expression(argument, current_sheet)?; + } + Ok(()) + } + SpreadsheetFormulaExpressionKind::Array { rows } => { + for value in rows.iter().flatten() { + self.collect_expression(value, current_sheet)?; + } + Ok(()) + } + SpreadsheetFormulaExpressionKind::Literal(_) + | SpreadsheetFormulaExpressionKind::Reference(_) => Ok(()), + } + } + + fn try_reference_areas( + &mut self, + expression: &SpreadsheetFormulaExpression, + current_sheet: usize, + ) -> UseResult>> { + match &expression.kind { + SpreadsheetFormulaExpressionKind::Reference(reference) => { + Ok(Some(self.single_reference_areas(reference, current_sheet)?)) + } + SpreadsheetFormulaExpressionKind::Binary { + operator: + SpreadsheetFormulaBinaryOperator::Range + | SpreadsheetFormulaBinaryOperator::Union + | SpreadsheetFormulaBinaryOperator::Intersection, + left, + right, + } => { + let SpreadsheetFormulaExpressionKind::Binary { operator, .. } = &expression.kind + else { + return Ok(None); + }; + if matches!(operator, SpreadsheetFormulaBinaryOperator::Range) { + return self.range_areas(left, right, current_sheet); + } + let Some(mut left) = self.try_reference_areas(left, current_sheet)? else { + return Ok(None); + }; + let Some(right) = self.try_reference_areas(right, current_sheet)? else { + return Ok(None); + }; + if matches!(operator, SpreadsheetFormulaBinaryOperator::Union) { + ensure_reference_area_total(left.len(), right.len())?; + left.extend(right); + Ok(Some(left)) + } else { + ensure_reference_comparisons(left.len(), right.len())?; + let mut areas = Vec::new(); + for left in left { + for right in &right { + if let Some(area) = left.intersect(*right) { + push_reference_area(&mut areas, area)?; + } + } + } + Ok(Some(areas)) + } + } + SpreadsheetFormulaExpressionKind::Parenthesized(inner) => { + self.try_reference_areas(inner, current_sheet) + } + SpreadsheetFormulaExpressionKind::Unary { + operator: SpreadsheetFormulaUnaryOperator::ImplicitIntersection, + operand, + } + | SpreadsheetFormulaExpressionKind::Postfix { + operator: SpreadsheetFormulaPostfixOperator::Spill, + operand, + } => self.try_reference_areas(operand, current_sheet), + _ => Ok(None), + } + } + + fn single_reference_areas( + &mut self, + reference: &SpreadsheetFormulaReference, + current_sheet: usize, + ) -> UseResult> { + let (start_column, start_row, end_column, end_row) = match reference.kind { + SpreadsheetFormulaReferenceKind::Cell { column, row, .. } => (column, row, column, row), + SpreadsheetFormulaReferenceKind::Column { column, .. } => (column, 1, column, MAX_ROWS), + SpreadsheetFormulaReferenceKind::Row { row, .. } => (1, row, MAX_COLUMNS, row), + }; + Ok(self + .resolve_sheets(reference.qualifier.as_ref(), current_sheet)? + .into_iter() + .map(|sheet_index| FormulaReferenceArea { + sheet_index, + start_column, + start_row, + end_column, + end_row, + }) + .collect()) + } + + fn range_areas( + &mut self, + left: &SpreadsheetFormulaExpression, + right: &SpreadsheetFormulaExpression, + current_sheet: usize, + ) -> UseResult>> { + let Some(left) = endpoint_reference(left) else { + return Ok(None); + }; + let Some(right) = endpoint_reference(right) else { + return Ok(None); + }; + let qualifier = match (&left.qualifier, &right.qualifier) { + (Some(left), Some(right)) if left != right => { + self.push_unresolved( + SpreadsheetFormulaUnresolvedReferenceKind::UnsupportedReference, + "range with different worksheet qualifiers", + ); + return Ok(Some(Vec::new())); + } + (Some(qualifier), _) | (_, Some(qualifier)) => Some(qualifier), + (None, None) => None, + }; + let coordinates = match (left.kind, right.kind) { + ( + SpreadsheetFormulaReferenceKind::Cell { + column: left_column, + row: left_row, + .. + }, + SpreadsheetFormulaReferenceKind::Cell { + column: right_column, + row: right_row, + .. + }, + ) => ( + left_column.min(right_column), + left_row.min(right_row), + left_column.max(right_column), + left_row.max(right_row), + ), + ( + SpreadsheetFormulaReferenceKind::Column { + column: left_column, + .. + }, + SpreadsheetFormulaReferenceKind::Column { + column: right_column, + .. + }, + ) => ( + left_column.min(right_column), + 1, + left_column.max(right_column), + MAX_ROWS, + ), + ( + SpreadsheetFormulaReferenceKind::Row { row: left_row, .. }, + SpreadsheetFormulaReferenceKind::Row { row: right_row, .. }, + ) => ( + 1, + left_row.min(right_row), + MAX_COLUMNS, + left_row.max(right_row), + ), + _ => return Ok(None), + }; + Ok(Some( + self.resolve_sheets(qualifier, current_sheet)? + .into_iter() + .map(|sheet_index| FormulaReferenceArea { + sheet_index, + start_column: coordinates.0, + start_row: coordinates.1, + end_column: coordinates.2, + end_row: coordinates.3, + }) + .collect(), + )) + } + + fn resolve_sheets( + &mut self, + qualifier: Option<&SpreadsheetFormulaQualifier>, + current_sheet: usize, + ) -> UseResult> { + let Some(qualifier) = qualifier else { + return Ok(vec![current_sheet]); + }; + if qualifier.is_external() { + self.push_unresolved( + SpreadsheetFormulaUnresolvedReferenceKind::ExternalWorkbook, + qualifier_label(qualifier), + ); + return Ok(Vec::new()); + } + let Some(start) = self.sheet_position(&qualifier.worksheet) else { + self.push_unresolved( + SpreadsheetFormulaUnresolvedReferenceKind::MissingWorksheet, + &qualifier.worksheet, + ); + return Ok(Vec::new()); + }; + let Some(end_name) = qualifier.worksheet_end.as_deref() else { + return Ok(vec![start]); + }; + let Some(end) = self.sheet_position(end_name) else { + self.push_unresolved( + SpreadsheetFormulaUnresolvedReferenceKind::MissingWorksheet, + end_name, + ); + return Ok(Vec::new()); + }; + let low = start.min(end); + let high = start.max(end); + let areas = high + .checked_sub(low) + .and_then(|distance| distance.checked_add(1)) + .ok_or_else(reference_area_limit)?; + if areas > MAX_SPREADSHEET_FORMULA_REFERENCE_AREAS { + return Err(reference_area_limit().with_detail("areas", areas)); + } + Ok((low..=high).collect()) + } + + fn collect_named_reference( + &mut self, + qualifier: Option<&SpreadsheetFormulaQualifier>, + name: &str, + current_sheet: usize, + ) -> UseResult<()> { + if qualifier.is_some_and(SpreadsheetFormulaQualifier::is_external) { + self.push_unresolved( + SpreadsheetFormulaUnresolvedReferenceKind::ExternalWorkbook, + qualifier.map_or(name.to_string(), qualifier_label), + ); + return Ok(()); + } + let explicit_scope = qualifier.and_then(|qualifier| { + if qualifier.is_three_dimensional() { + None + } else { + self.sheet_position(&qualifier.worksheet) + } + }); + if qualifier.is_some() && explicit_scope.is_none() { + if let Some(qualifier) = qualifier { + let kind = if self.sheet_position(&qualifier.worksheet).is_none() { + SpreadsheetFormulaUnresolvedReferenceKind::MissingWorksheet + } else { + SpreadsheetFormulaUnresolvedReferenceKind::UnsupportedReference + }; + self.push_unresolved(kind, qualifier_label(qualifier)); + } + return Ok(()); + } + let normalized = name.to_lowercase(); + let local_scope = explicit_scope.or_else(|| qualifier.is_none().then_some(current_sheet)); + let definition = local_scope + .and_then(|scope| { + self.named_definitions + .get(&(Some(scope), normalized.clone())) + }) + .or_else(|| self.named_definitions.get(&(None, normalized.clone()))) + .cloned(); + let Some(definition) = definition else { + self.push_unresolved( + SpreadsheetFormulaUnresolvedReferenceKind::UndefinedName, + name, + ); + return Ok(()); + }; + if self.visited_names.len() >= MAX_SPREADSHEET_FORMULA_DEPTH { + self.push_unresolved( + SpreadsheetFormulaUnresolvedReferenceKind::NamedRangeDepth, + &definition.name, + ); + return Ok(()); + } + let key = (definition.scope_sheet, normalized); + if !self.visited_names.insert(key.clone()) { + self.push_unresolved( + SpreadsheetFormulaUnresolvedReferenceKind::NamedRangeCycle, + &definition.name, + ); + return Ok(()); + } + let parsed = parse_spreadsheet_formula(&definition.formula).map_err(|error| { + error + .with_detail("namedRange", definition.name.clone()) + .with_detail( + "scope", + definition.scope_sheet.map_or_else( + || "workbook".to_string(), + |scope| { + self.sheet_names + .get(scope) + .cloned() + .unwrap_or_else(|| "unknown".to_string()) + }, + ), + ) + })?; + let definition_sheet = definition.scope_sheet.unwrap_or(current_sheet); + let result = self.collect_expression(&parsed.root, definition_sheet); + self.visited_names.remove(&key); + result + } + + fn sheet_position(&self, name: &str) -> Option { + self.sheet_names + .iter() + .position(|sheet| sheet.eq_ignore_ascii_case(name)) + } + + fn push_unresolved( + &mut self, + kind: SpreadsheetFormulaUnresolvedReferenceKind, + reference: impl AsRef, + ) { + self.unresolved.push(SpreadsheetFormulaUnresolvedReference { + kind, + reference: reference.as_ref().to_string(), + }); + } + + fn extend_areas(&mut self, areas: Vec) -> UseResult<()> { + ensure_reference_area_total(self.areas.len(), areas.len())?; + self.areas.extend(areas); + Ok(()) + } +} + +fn endpoint_reference( + expression: &SpreadsheetFormulaExpression, +) -> Option<&SpreadsheetFormulaReference> { + match &expression.kind { + SpreadsheetFormulaExpressionKind::Reference(reference) => Some(reference), + SpreadsheetFormulaExpressionKind::Parenthesized(inner) => endpoint_reference(inner), + _ => None, + } +} + +fn qualifier_label(qualifier: &SpreadsheetFormulaQualifier) -> String { + let workbook = qualifier.workbook.as_deref().unwrap_or_default(); + let worksheets = qualifier.worksheet_end.as_ref().map_or_else( + || qualifier.worksheet.clone(), + |end| format!("{}:{end}", qualifier.worksheet), + ); + format!("{workbook}{worksheets}") +} + +fn push_reference_area( + areas: &mut Vec, + area: FormulaReferenceArea, +) -> UseResult<()> { + let total = areas + .len() + .checked_add(1) + .ok_or_else(reference_area_limit)?; + if total > MAX_SPREADSHEET_FORMULA_REFERENCE_AREAS { + return Err(reference_area_limit().with_detail("areas", total)); + } + areas.push(area); + Ok(()) +} + +fn ensure_reference_area_total(left: usize, right: usize) -> UseResult<()> { + let total = left.checked_add(right).ok_or_else(reference_area_limit)?; + if total > MAX_SPREADSHEET_FORMULA_REFERENCE_AREAS { + return Err(reference_area_limit().with_detail("areas", total)); + } + Ok(()) +} + +fn ensure_reference_comparisons(left: usize, right: usize) -> UseResult<()> { + let visits = left.checked_mul(right).ok_or_else(reference_visit_limit)?; + if visits > MAX_SPREADSHEET_FORMULA_REFERENCE_VISITS { + return Err(reference_visit_limit().with_detail("visits", visits)); + } + Ok(()) +} + +fn reference_area_limit() -> a3s_use_core::UseError { + graph_error( + "use.office.spreadsheet_formula_reference_area_limit", + format!( + "Spreadsheet formulas retain at most {MAX_SPREADSHEET_FORMULA_REFERENCE_AREAS} static reference areas." + ), + ) +} + +fn reference_visit_limit() -> a3s_use_core::UseError { + graph_error( + "use.office.spreadsheet_formula_reference_visit_limit", + format!( + "Spreadsheet formula reference operators visit at most {MAX_SPREADSHEET_FORMULA_REFERENCE_VISITS} area pairs." + ), + ) +} diff --git a/crates/office/src/spreadsheet_formula/lexer.rs b/crates/office/src/spreadsheet_formula/lexer.rs new file mode 100644 index 00000000..d74e2142 --- /dev/null +++ b/crates/office/src/spreadsheet_formula/lexer.rs @@ -0,0 +1,823 @@ +use crate::spreadsheet_reference::{MAX_COLUMNS, MAX_ROWS}; + +use super::{ + ast::{ + SpreadsheetFormulaErrorLiteral, SpreadsheetFormulaQualifier, SpreadsheetFormulaReference, + SpreadsheetFormulaReferenceKind, SpreadsheetFormulaSpan, MAX_SPREADSHEET_FORMULA_NODES, + }, + FormulaParseFailure, +}; + +const MAX_LEXICAL_TOKENS: usize = MAX_SPREADSHEET_FORMULA_NODES * 2 + 1; + +#[derive(Debug, Clone, PartialEq, Eq)] +pub(super) struct FormulaToken { + pub(super) kind: FormulaTokenKind, + pub(super) span: SpreadsheetFormulaSpan, +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub(super) enum FormulaTokenKind { + Number(String), + Text(String), + Error(SpreadsheetFormulaErrorLiteral), + Reference(SpreadsheetFormulaReference), + Name { + qualifier: Option, + name: String, + }, + StructuredReference { + qualifier: Option, + reference: String, + }, + Space, + Plus, + Minus, + Star, + Slash, + Caret, + Ampersand, + Equal, + NotEqual, + LessThan, + LessThanOrEqual, + GreaterThan, + GreaterThanOrEqual, + Percent, + Hash, + At, + Colon, + Comma, + Semicolon, + LeftParen, + RightParen, + LeftBrace, + RightBrace, + End, +} + +impl FormulaTokenKind { + pub(super) const fn description(&self) -> &'static str { + match self { + Self::Number(_) => "number", + Self::Text(_) => "string", + Self::Error(_) => "error literal", + Self::Reference(_) => "cell reference", + Self::Name { .. } => "name", + Self::StructuredReference { .. } => "structured reference", + Self::Space => "space", + Self::Plus => "`+`", + Self::Minus => "`-`", + Self::Star => "`*`", + Self::Slash => "`/`", + Self::Caret => "`^`", + Self::Ampersand => "`&`", + Self::Equal => "`=`", + Self::NotEqual => "`<>`", + Self::LessThan => "`<`", + Self::LessThanOrEqual => "`<=`", + Self::GreaterThan => "`>`", + Self::GreaterThanOrEqual => "`>=`", + Self::Percent => "`%`", + Self::Hash => "`#`", + Self::At => "`@`", + Self::Colon => "`:`", + Self::Comma => "`,`", + Self::Semicolon => "`;`", + Self::LeftParen => "`(`", + Self::RightParen => "`)`", + Self::LeftBrace => "`{`", + Self::RightBrace => "`}`", + Self::End => "end of formula", + } + } +} + +pub(super) fn lex(source: &str) -> Result, FormulaParseFailure> { + FormulaLexer::new(source).lex() +} + +struct FormulaLexer<'a> { + source: &'a str, + cursor: usize, + tokens: Vec, +} + +impl<'a> FormulaLexer<'a> { + fn new(source: &'a str) -> Self { + Self { + source, + cursor: 0, + tokens: Vec::with_capacity(source.len().min(MAX_LEXICAL_TOKENS)), + } + } + + fn lex(mut self) -> Result, FormulaParseFailure> { + while self.cursor < self.source.len() { + let start = self.cursor; + let character = self.current_character().ok_or_else(|| { + FormulaParseFailure::new(self.cursor, "Formula contains invalid UTF-8 boundaries.") + })?; + if character.is_whitespace() { + self.consume_while(char::is_whitespace); + self.push(FormulaTokenKind::Space, start, self.cursor)?; + continue; + } + if character == '"' { + let token = self.lex_text()?; + self.push_token(token)?; + continue; + } + if character == '\'' { + let token = self.lex_qualified_atom()?.ok_or_else(|| { + FormulaParseFailure::new(start, "Worksheet qualifier is invalid.") + })?; + self.push_token(token)?; + continue; + } + if let Some(token) = self.lex_qualified_atom()? { + self.push_token(token)?; + continue; + } + if character == '$' || character.is_ascii_alphabetic() { + match self.try_cell_reference(self.cursor, None, false)? { + Some((reference, end)) => { + self.cursor = end; + self.push(FormulaTokenKind::Reference(reference), start, self.cursor)?; + continue; + } + None if character == '$' => { + let (reference, end) = self.lex_absolute_axis(None)?; + self.cursor = end; + self.push(FormulaTokenKind::Reference(reference), start, self.cursor)?; + continue; + } + None => {} + } + } + if character.is_ascii_digit() + || (character == '.' + && self + .next_character() + .is_some_and(|value| value.is_ascii_digit())) + { + let token = self.lex_number()?; + self.push_token(token)?; + continue; + } + if is_name_start(character) || character == '[' { + let token = self.lex_name(None, start)?; + self.push_token(token)?; + continue; + } + if character == '#' { + if let Some(token) = self.lex_error_literal()? { + self.push_token(token)?; + } else { + self.cursor += character.len_utf8(); + self.push(FormulaTokenKind::Hash, start, self.cursor)?; + } + continue; + } + + self.cursor += character.len_utf8(); + let kind = match character { + '+' => FormulaTokenKind::Plus, + '-' => FormulaTokenKind::Minus, + '*' => FormulaTokenKind::Star, + '/' => FormulaTokenKind::Slash, + '^' => FormulaTokenKind::Caret, + '&' => FormulaTokenKind::Ampersand, + '=' => FormulaTokenKind::Equal, + '%' => FormulaTokenKind::Percent, + '@' => FormulaTokenKind::At, + ':' => FormulaTokenKind::Colon, + ',' => FormulaTokenKind::Comma, + ';' => FormulaTokenKind::Semicolon, + '(' => FormulaTokenKind::LeftParen, + ')' => FormulaTokenKind::RightParen, + '{' => FormulaTokenKind::LeftBrace, + '}' => FormulaTokenKind::RightBrace, + '<' => { + if self.source[self.cursor..].starts_with('=') { + self.cursor += 1; + FormulaTokenKind::LessThanOrEqual + } else if self.source[self.cursor..].starts_with('>') { + self.cursor += 1; + FormulaTokenKind::NotEqual + } else { + FormulaTokenKind::LessThan + } + } + '>' => { + if self.source[self.cursor..].starts_with('=') { + self.cursor += 1; + FormulaTokenKind::GreaterThanOrEqual + } else { + FormulaTokenKind::GreaterThan + } + } + _ => { + return Err(FormulaParseFailure::new( + start, + format!("Unexpected character `{character}`."), + )); + } + }; + self.push(kind, start, self.cursor)?; + } + self.push(FormulaTokenKind::End, self.source.len(), self.source.len())?; + Ok(self.tokens) + } + + fn lex_text(&mut self) -> Result { + let start = self.cursor; + self.cursor += 1; + let mut value = String::new(); + while self.cursor < self.source.len() { + let character = self.current_character().ok_or_else(|| { + FormulaParseFailure::new(self.cursor, "String literal has an invalid boundary.") + })?; + if character == '"' { + self.cursor += 1; + if self.source[self.cursor..].starts_with('"') { + value.push('"'); + self.cursor += 1; + continue; + } + return Ok(FormulaToken { + kind: FormulaTokenKind::Text(value), + span: SpreadsheetFormulaSpan::new(start, self.cursor), + }); + } + value.push(character); + self.cursor += character.len_utf8(); + } + Err(FormulaParseFailure::new( + self.source.len(), + "String literal is not closed.", + )) + } + + fn lex_qualified_atom(&mut self) -> Result, FormulaParseFailure> { + let start = self.cursor; + let Some((qualifier, target_start)) = self.scan_qualifier()? else { + return Ok(None); + }; + self.cursor = target_start; + if self.cursor >= self.source.len() { + return Err(FormulaParseFailure::new( + self.cursor, + "Worksheet qualifier must be followed by a reference or name.", + )); + } + if let Some((reference, end)) = + self.try_cell_reference(self.cursor, Some(qualifier.clone()), true)? + { + self.cursor = end; + return Ok(Some(FormulaToken { + kind: FormulaTokenKind::Reference(reference), + span: SpreadsheetFormulaSpan::new(start, end), + })); + } + if self.source[self.cursor..].starts_with('$') { + let (reference, end) = self.lex_absolute_axis(Some(qualifier))?; + self.cursor = end; + return Ok(Some(FormulaToken { + kind: FormulaTokenKind::Reference(reference), + span: SpreadsheetFormulaSpan::new(start, end), + })); + } + let token = self.lex_name(Some(qualifier), start)?; + Ok(Some(token)) + } + + fn scan_qualifier( + &self, + ) -> Result, FormulaParseFailure> { + let start = self.cursor; + let character = self.source[start..].chars().next().ok_or_else(|| { + FormulaParseFailure::new(start, "Worksheet qualifier has an invalid boundary.") + })?; + if character == '\'' { + let mut cursor = start + 1; + let mut decoded = String::new(); + while cursor < self.source.len() { + let current = self.source[cursor..].chars().next().ok_or_else(|| { + FormulaParseFailure::new(cursor, "Worksheet qualifier is not valid UTF-8.") + })?; + if current == '\'' { + let after = cursor + 1; + if self.source[after..].starts_with('\'') { + decoded.push('\''); + cursor = after + 1; + continue; + } + if self.source[after..].starts_with('!') { + let qualifier = parse_qualifier_value(decoded, start)?; + return Ok(Some((qualifier, after + 1))); + } + return Err(FormulaParseFailure::new( + after, + "A quoted worksheet qualifier must be followed by `!`.", + )); + } + decoded.push(current); + cursor += current.len_utf8(); + } + return Err(FormulaParseFailure::new( + self.source.len(), + "Quoted worksheet qualifier is not closed.", + )); + } + if !is_bare_qualifier_character(character) { + return Ok(None); + } + + let mut cursor = start; + while cursor < self.source.len() { + let current = self.source[cursor..].chars().next().ok_or_else(|| { + FormulaParseFailure::new(cursor, "Worksheet qualifier has an invalid boundary.") + })?; + if current == '!' { + if cursor == start { + return Ok(None); + } + let qualifier = + parse_qualifier_value(self.source[start..cursor].to_string(), start)?; + return Ok(Some((qualifier, cursor + 1))); + } + if !is_bare_qualifier_character(current) { + return Ok(None); + } + cursor += current.len_utf8(); + } + Ok(None) + } + + fn try_cell_reference( + &self, + start: usize, + qualifier: Option, + qualified: bool, + ) -> Result, FormulaParseFailure> { + let bytes = self.source.as_bytes(); + let mut cursor = start; + let absolute_column = bytes.get(cursor) == Some(&b'$'); + cursor += usize::from(absolute_column); + let column_start = cursor; + while bytes.get(cursor).is_some_and(u8::is_ascii_alphabetic) { + cursor += 1; + } + let column_length = cursor.saturating_sub(column_start); + if column_length == 0 || column_length > 3 { + return Ok(None); + } + let absolute_row = bytes.get(cursor) == Some(&b'$'); + cursor += usize::from(absolute_row); + let row_start = cursor; + while bytes.get(cursor).is_some_and(u8::is_ascii_digit) { + cursor += 1; + } + if cursor == row_start { + return Ok(None); + } + if !qualified && bytes.get(cursor) == Some(&b'(') { + return Ok(None); + } + if self + .source + .get(cursor..) + .and_then(|value| value.chars().next()) + .is_some_and(is_name_continue) + { + return Ok(None); + } + let column = parse_ascii_column(&self.source[column_start..column_start + column_length]) + .filter(|value| *value <= MAX_COLUMNS) + .ok_or_else(|| { + FormulaParseFailure::new(column_start, "Cell reference column is outside A:XFD.") + })?; + let row = self.source[row_start..cursor] + .parse::() + .ok() + .filter(|value| (1..=MAX_ROWS).contains(value)) + .ok_or_else(|| { + FormulaParseFailure::new(row_start, "Cell reference row is outside 1:1048576.") + })?; + Ok(Some(( + SpreadsheetFormulaReference { + qualifier, + kind: SpreadsheetFormulaReferenceKind::Cell { + column, + row, + absolute_column, + absolute_row, + }, + }, + cursor, + ))) + } + + fn lex_absolute_axis( + &self, + qualifier: Option, + ) -> Result<(SpreadsheetFormulaReference, usize), FormulaParseFailure> { + let start = self.cursor; + let after_dollar = start + 1; + let Some(character) = self.source[after_dollar..].chars().next() else { + return Err(FormulaParseFailure::new( + start, + "Absolute reference marker must be followed by a row or column.", + )); + }; + if character.is_ascii_alphabetic() { + let mut cursor = after_dollar; + while self + .source + .as_bytes() + .get(cursor) + .is_some_and(u8::is_ascii_alphabetic) + { + cursor += 1; + } + if self + .source + .get(cursor..) + .and_then(|value| value.chars().next()) + .is_some_and(is_name_continue) + { + return Err(FormulaParseFailure::new( + cursor, + "Absolute column reference has invalid trailing characters.", + )); + } + let column = parse_ascii_column(&self.source[after_dollar..cursor]) + .filter(|value| *value <= MAX_COLUMNS) + .ok_or_else(|| { + FormulaParseFailure::new(after_dollar, "Column reference is outside A:XFD.") + })?; + return Ok(( + SpreadsheetFormulaReference { + qualifier, + kind: SpreadsheetFormulaReferenceKind::Column { + column, + absolute: true, + }, + }, + cursor, + )); + } + if character.is_ascii_digit() { + let mut cursor = after_dollar; + while self + .source + .as_bytes() + .get(cursor) + .is_some_and(u8::is_ascii_digit) + { + cursor += 1; + } + let row = self.source[after_dollar..cursor] + .parse::() + .ok() + .filter(|value| (1..=MAX_ROWS).contains(value)) + .ok_or_else(|| { + FormulaParseFailure::new(after_dollar, "Row reference is outside 1:1048576.") + })?; + return Ok(( + SpreadsheetFormulaReference { + qualifier, + kind: SpreadsheetFormulaReferenceKind::Row { + row, + absolute: true, + }, + }, + cursor, + )); + } + Err(FormulaParseFailure::new( + after_dollar, + "Absolute reference marker must be followed by an ASCII row or column.", + )) + } + + fn lex_number(&mut self) -> Result { + let start = self.cursor; + let bytes = self.source.as_bytes(); + while bytes.get(self.cursor).is_some_and(u8::is_ascii_digit) { + self.cursor += 1; + } + if bytes.get(self.cursor) == Some(&b'.') { + self.cursor += 1; + while bytes.get(self.cursor).is_some_and(u8::is_ascii_digit) { + self.cursor += 1; + } + } + if matches!(bytes.get(self.cursor), Some(b'e' | b'E')) { + let exponent = self.cursor; + self.cursor += 1; + if matches!(bytes.get(self.cursor), Some(b'+' | b'-')) { + self.cursor += 1; + } + let digits = self.cursor; + while bytes.get(self.cursor).is_some_and(u8::is_ascii_digit) { + self.cursor += 1; + } + if self.cursor == digits { + return Err(FormulaParseFailure::new( + exponent, + "Numeric exponent has no digits.", + )); + } + } + if self + .source + .get(self.cursor..) + .and_then(|value| value.chars().next()) + .is_some_and(is_name_continue) + { + return Err(FormulaParseFailure::new( + self.cursor, + "Numeric literal has invalid trailing characters.", + )); + } + let raw = &self.source[start..self.cursor]; + if !raw.parse::().ok().is_some_and(f64::is_finite) { + return Err(FormulaParseFailure::new( + start, + "Numeric literal must be finite.", + )); + } + Ok(FormulaToken { + kind: FormulaTokenKind::Number(raw.to_string()), + span: SpreadsheetFormulaSpan::new(start, self.cursor), + }) + } + + fn lex_name( + &mut self, + qualifier: Option, + span_start: usize, + ) -> Result { + let name_start = self.cursor; + if self.source[self.cursor..].starts_with('[') { + self.scan_bracket_groups()?; + } else { + let Some(first) = self.current_character() else { + return Err(FormulaParseFailure::new( + self.cursor, + "Expected a formula name.", + )); + }; + if !(is_name_start(first) || (qualifier.is_some() && first.is_ascii_digit())) { + return Err(FormulaParseFailure::new( + self.cursor, + "Worksheet qualifier must be followed by a reference or name.", + )); + } + self.cursor += first.len_utf8(); + self.consume_while(|character| { + is_name_continue(character) && !matches!(character, '[' | ']') + }); + while self.source[self.cursor..].starts_with('[') { + self.scan_bracket_groups()?; + } + } + let name = self.source[name_start..self.cursor].to_string(); + if name.is_empty() { + return Err(FormulaParseFailure::new( + name_start, + "Formula name is empty.", + )); + } + let kind = if name.contains('[') { + FormulaTokenKind::StructuredReference { + qualifier, + reference: name, + } + } else { + FormulaTokenKind::Name { qualifier, name } + }; + Ok(FormulaToken { + kind, + span: SpreadsheetFormulaSpan::new(span_start, self.cursor), + }) + } + + fn scan_bracket_groups(&mut self) -> Result<(), FormulaParseFailure> { + while self.source[self.cursor..].starts_with('[') { + let start = self.cursor; + let mut depth = 0_usize; + while self.cursor < self.source.len() { + let character = self.current_character().ok_or_else(|| { + FormulaParseFailure::new( + self.cursor, + "Structured reference has an invalid boundary.", + ) + })?; + self.cursor += character.len_utf8(); + if character == '\'' { + let escaped = self.current_character().ok_or_else(|| { + FormulaParseFailure::new( + self.cursor, + "Structured reference escape has no following character.", + ) + })?; + self.cursor += escaped.len_utf8(); + continue; + } + match character { + '[' => depth = depth.saturating_add(1), + ']' => { + depth = depth.saturating_sub(1); + if depth == 0 { + break; + } + } + _ => {} + } + } + if depth != 0 { + return Err(FormulaParseFailure::new( + start, + "Structured reference bracket is not closed.", + )); + } + } + Ok(()) + } + + fn lex_error_literal(&mut self) -> Result, FormulaParseFailure> { + const ERRORS: &[SpreadsheetFormulaErrorLiteral] = &[ + SpreadsheetFormulaErrorLiteral::GettingData, + SpreadsheetFormulaErrorLiteral::DivisionByZero, + SpreadsheetFormulaErrorLiteral::NotAvailable, + SpreadsheetFormulaErrorLiteral::Calculation, + SpreadsheetFormulaErrorLiteral::Reference, + SpreadsheetFormulaErrorLiteral::Blocked, + SpreadsheetFormulaErrorLiteral::Unknown, + SpreadsheetFormulaErrorLiteral::Connect, + SpreadsheetFormulaErrorLiteral::Python, + SpreadsheetFormulaErrorLiteral::Value, + SpreadsheetFormulaErrorLiteral::Field, + SpreadsheetFormulaErrorLiteral::Spill, + SpreadsheetFormulaErrorLiteral::Number, + SpreadsheetFormulaErrorLiteral::Name, + SpreadsheetFormulaErrorLiteral::Busy, + SpreadsheetFormulaErrorLiteral::Null, + ]; + let start = self.cursor; + for literal in ERRORS { + let raw = literal.as_str(); + let Some(candidate) = self.source.get(start..start.saturating_add(raw.len())) else { + continue; + }; + if candidate.eq_ignore_ascii_case(raw) { + self.cursor += raw.len(); + return Ok(Some(FormulaToken { + kind: FormulaTokenKind::Error(*literal), + span: SpreadsheetFormulaSpan::new(start, self.cursor), + })); + } + } + Ok(None) + } + + fn current_character(&self) -> Option { + self.source.get(self.cursor..)?.chars().next() + } + + fn next_character(&self) -> Option { + let current = self.current_character()?; + self.source + .get(self.cursor + current.len_utf8()..)? + .chars() + .next() + } + + fn consume_while(&mut self, predicate: impl Fn(char) -> bool) { + while let Some(character) = self.current_character() { + if !predicate(character) { + break; + } + self.cursor += character.len_utf8(); + } + } + + fn push_token(&mut self, token: FormulaToken) -> Result<(), FormulaParseFailure> { + if self.tokens.len() >= MAX_LEXICAL_TOKENS { + return Err(FormulaParseFailure::new( + token.span.start, + "Formula contains too many lexical tokens.", + )); + } + self.tokens.push(token); + Ok(()) + } + + fn push( + &mut self, + kind: FormulaTokenKind, + start: usize, + end: usize, + ) -> Result<(), FormulaParseFailure> { + self.push_token(FormulaToken { + kind, + span: SpreadsheetFormulaSpan::new(start, end), + }) + } +} + +fn parse_qualifier_value( + value: String, + position: usize, +) -> Result { + if value.is_empty() { + return Err(FormulaParseFailure::new( + position, + "Worksheet qualifier is empty.", + )); + } + let (workbook, worksheets) = value.rfind(']').map_or((None, value.as_str()), |end| { + ( + Some(value[..=end].to_string()), + value.get(end + 1..).unwrap_or_default(), + ) + }); + if worksheets.is_empty() { + return Err(FormulaParseFailure::new( + position, + "Worksheet qualifier has no worksheet name.", + )); + } + let (worksheet, worksheet_end) = worksheets + .split_once(':') + .map_or((worksheets, None), |(start, end)| (start, Some(end))); + if worksheet.is_empty() || worksheet_end.is_some_and(str::is_empty) { + return Err(FormulaParseFailure::new( + position, + "Three-dimensional worksheet qualifier has an empty endpoint.", + )); + } + Ok(SpreadsheetFormulaQualifier { + workbook, + worksheet: worksheet.to_string(), + worksheet_end: worksheet_end.map(ToOwned::to_owned), + }) +} + +fn parse_ascii_column(value: &str) -> Option { + value.bytes().try_fold(0_u32, |column, byte| { + column.checked_mul(26).and_then(|column| { + column.checked_add(u32::from(byte.to_ascii_uppercase().checked_sub(b'A')?) + 1) + }) + }) +} + +fn is_bare_qualifier_character(character: char) -> bool { + character.is_alphanumeric() || matches!(character, '_' | '.' | '\\' | '[' | ']' | ':' | '$') +} + +fn is_name_start(character: char) -> bool { + character.is_alphabetic() || matches!(character, '_' | '\\' | '?') +} + +fn is_name_continue(character: char) -> bool { + character.is_alphanumeric() || matches!(character, '_' | '.' | '\\' | '?') +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn lexer_decodes_literals_and_qualified_references_with_utf8_spans() { + let source = r#""a""b"+'销售 数据'!$C$3+[Book.xlsx]Data!A1"#; + let tokens = lex(source).unwrap(); + assert_eq!(tokens[0].kind, FormulaTokenKind::Text("a\"b".to_string())); + assert_eq!( + &source[tokens[0].span.start..tokens[0].span.end], + r#""a""b""# + ); + assert!(matches!(tokens[2].kind, FormulaTokenKind::Reference(_))); + assert!(matches!(tokens[4].kind, FormulaTokenKind::Reference(_))); + let FormulaTokenKind::Reference(reference) = &tokens[2].kind else { + unreachable!(); + }; + let qualifier = reference.qualifier.as_ref().unwrap(); + assert_eq!(qualifier.worksheet, "销售 数据"); + assert!(!qualifier.is_external()); + let FormulaTokenKind::Reference(reference) = &tokens[4].kind else { + unreachable!(); + }; + assert!(reference.qualifier.as_ref().unwrap().is_external()); + } + + #[test] + fn lexer_rejects_unclosed_strings_qualifiers_and_structured_references() { + for source in ["\"open", "'Sheet 1!A1", "Table1[[Column]"] { + assert!(lex(source).is_err(), "{source}"); + } + } +} diff --git a/crates/office/src/spreadsheet_formula/parser.rs b/crates/office/src/spreadsheet_formula/parser.rs new file mode 100644 index 00000000..e8d7b69d --- /dev/null +++ b/crates/office/src/spreadsheet_formula/parser.rs @@ -0,0 +1,923 @@ +use crate::spreadsheet_reference::{MAX_COLUMNS, MAX_ROWS}; + +use super::{ + ast::{ + SpreadsheetFormula, SpreadsheetFormulaBinaryOperator, SpreadsheetFormulaExpression, + SpreadsheetFormulaExpressionKind, SpreadsheetFormulaLiteral, + SpreadsheetFormulaPostfixOperator, SpreadsheetFormulaReference, + SpreadsheetFormulaReferenceKind, SpreadsheetFormulaSpan, SpreadsheetFormulaUnaryOperator, + MAX_SPREADSHEET_FORMULA_DEPTH, MAX_SPREADSHEET_FORMULA_NODES, + }, + lexer::{FormulaToken, FormulaTokenKind}, + FormulaParseFailure, +}; + +const MAX_FUNCTION_ARGUMENTS: usize = 255; + +const BINDING_COMPARISON: u8 = 10; +const BINDING_CONCATENATE: u8 = 20; +const BINDING_ADDITIVE: u8 = 30; +const BINDING_MULTIPLICATIVE: u8 = 40; +const BINDING_POWER: u8 = 50; +const BINDING_POSTFIX: u8 = 60; +const BINDING_PREFIX: u8 = 70; +const BINDING_UNION: u8 = 80; +const BINDING_INTERSECTION: u8 = 90; +const BINDING_RANGE: u8 = 100; + +pub(super) fn parse( + source: &str, + tokens: Vec, +) -> Result { + FormulaParser::new(source, tokens).parse() +} + +#[derive(Debug)] +struct ParsedExpression { + expression: SpreadsheetFormulaExpression, + depth: usize, + reference_like: bool, +} + +impl ParsedExpression { + fn span(&self) -> SpreadsheetFormulaSpan { + self.expression.span + } +} + +struct FormulaParser<'a> { + source: &'a str, + tokens: Vec, + cursor: usize, + nodes: usize, +} + +impl<'a> FormulaParser<'a> { + fn new(source: &'a str, tokens: Vec) -> Self { + Self { + source, + tokens, + cursor: 0, + nodes: 0, + } + } + + fn parse(mut self) -> Result { + let expression = self.parse_expression(0, true, 1)?; + self.skip_spaces(); + let trailing = self.current()?.clone(); + if !matches!(trailing.kind, FormulaTokenKind::End) { + return Err(FormulaParseFailure::new( + trailing.span.start, + format!( + "Unexpected {}; expected end of formula.", + trailing.kind.description() + ), + )); + } + validate_axis_references(&expression.expression, false)?; + Ok(SpreadsheetFormula { + root: expression.expression, + }) + } + + fn parse_expression( + &mut self, + minimum_binding: u8, + allow_union: bool, + call_depth: usize, + ) -> Result { + self.check_call_depth(call_depth)?; + let mut left = self.parse_prefix(allow_union, call_depth)?; + loop { + let space = self.consume_spaces(); + let token = self.current()?.clone(); + + if matches!( + token.kind, + FormulaTokenKind::Percent | FormulaTokenKind::Hash + ) && BINDING_POSTFIX >= minimum_binding + { + self.cursor += 1; + let operator = match token.kind { + FormulaTokenKind::Percent => SpreadsheetFormulaPostfixOperator::Percent, + FormulaTokenKind::Hash => SpreadsheetFormulaPostfixOperator::Spill, + _ => { + return Err(FormulaParseFailure::new( + token.span.start, + "Unsupported postfix operator.", + )); + } + }; + if matches!(operator, SpreadsheetFormulaPostfixOperator::Spill) + && !left.reference_like + { + return Err(FormulaParseFailure::new( + token.span.start, + "The spill operator requires a reference expression.", + )); + } + let span = left.span().through(token.span); + let depth = left.depth.saturating_add(1); + let reference_like = matches!(operator, SpreadsheetFormulaPostfixOperator::Spill); + left = self.node( + SpreadsheetFormulaExpressionKind::Postfix { + operator, + operand: Box::new(left.expression), + }, + span, + depth, + reference_like, + )?; + continue; + } + + let intersection = space + .filter(|_| left.reference_like && token_starts_reference_expression(&token.kind)); + let Some((operator, left_binding, right_binding, consume_token)) = + infix_operator(&token.kind, allow_union, intersection.is_some()) + else { + break; + }; + if left_binding < minimum_binding { + break; + } + let operator_span = if let Some(space) = intersection { + space + } else { + if consume_token { + self.cursor += 1; + } + token.span + }; + let right = + self.parse_expression(right_binding, allow_union, call_depth.saturating_add(1))?; + left = self.binary(left, right, operator, operator_span)?; + } + Ok(left) + } + + fn parse_prefix( + &mut self, + allow_union: bool, + call_depth: usize, + ) -> Result { + self.skip_spaces(); + let token = self.current()?.clone(); + self.cursor += 1; + match token.kind { + FormulaTokenKind::Number(value) => self.node( + SpreadsheetFormulaExpressionKind::Literal(SpreadsheetFormulaLiteral::Number(value)), + token.span, + 1, + false, + ), + FormulaTokenKind::Text(value) => self.node( + SpreadsheetFormulaExpressionKind::Literal(SpreadsheetFormulaLiteral::Text(value)), + token.span, + 1, + false, + ), + FormulaTokenKind::Error(value) => self.node( + SpreadsheetFormulaExpressionKind::Literal(SpreadsheetFormulaLiteral::Error(value)), + token.span, + 1, + false, + ), + FormulaTokenKind::Reference(reference) => self.node( + SpreadsheetFormulaExpressionKind::Reference(reference), + token.span, + 1, + true, + ), + FormulaTokenKind::StructuredReference { + qualifier, + reference, + } => self.node( + SpreadsheetFormulaExpressionKind::StructuredReference { + qualifier, + reference, + }, + token.span, + 1, + true, + ), + FormulaTokenKind::Name { qualifier, name } => { + if matches!( + self.current().map(|value| &value.kind), + Ok(FormulaTokenKind::LeftParen) + ) { + self.parse_function(qualifier, name, token.span, call_depth) + } else if qualifier.is_none() && name.eq_ignore_ascii_case("TRUE") { + self.node( + SpreadsheetFormulaExpressionKind::Literal( + SpreadsheetFormulaLiteral::Boolean(true), + ), + token.span, + 1, + false, + ) + } else if qualifier.is_none() && name.eq_ignore_ascii_case("FALSE") { + self.node( + SpreadsheetFormulaExpressionKind::Literal( + SpreadsheetFormulaLiteral::Boolean(false), + ), + token.span, + 1, + false, + ) + } else { + self.node( + SpreadsheetFormulaExpressionKind::Name { qualifier, name }, + token.span, + 1, + true, + ) + } + } + FormulaTokenKind::Plus | FormulaTokenKind::Minus | FormulaTokenKind::At => { + let operator = match token.kind { + FormulaTokenKind::Plus => SpreadsheetFormulaUnaryOperator::Positive, + FormulaTokenKind::Minus => SpreadsheetFormulaUnaryOperator::Negative, + FormulaTokenKind::At => SpreadsheetFormulaUnaryOperator::ImplicitIntersection, + _ => { + return Err(FormulaParseFailure::new( + token.span.start, + "Unsupported prefix operator.", + )); + } + }; + let operand = self.parse_expression( + BINDING_PREFIX, + allow_union, + call_depth.saturating_add(1), + )?; + if matches!( + operator, + SpreadsheetFormulaUnaryOperator::ImplicitIntersection + ) && !operand.reference_like + { + return Err(FormulaParseFailure::new( + token.span.start, + "Implicit intersection requires a reference expression.", + )); + } + let span = token.span.through(operand.span()); + let depth = operand.depth.saturating_add(1); + let reference_like = matches!( + operator, + SpreadsheetFormulaUnaryOperator::ImplicitIntersection + ); + self.node( + SpreadsheetFormulaExpressionKind::Unary { + operator, + operand: Box::new(operand.expression), + }, + span, + depth, + reference_like, + ) + } + FormulaTokenKind::LeftParen => { + let inner = self.parse_expression(0, true, call_depth.saturating_add(1))?; + self.skip_spaces(); + let close = self.expect(FormulaTokenKind::RightParen, "`)`")?; + let span = token.span.through(close.span); + let depth = inner.depth.saturating_add(1); + let reference_like = inner.reference_like; + self.node( + SpreadsheetFormulaExpressionKind::Parenthesized(Box::new(inner.expression)), + span, + depth, + reference_like, + ) + } + FormulaTokenKind::LeftBrace => self.parse_array(token.span, call_depth), + other => Err(FormulaParseFailure::new( + token.span.start, + format!("Expected an expression but found {}.", other.description()), + )), + } + } + + fn parse_function( + &mut self, + qualifier: Option, + name: String, + name_span: SpreadsheetFormulaSpan, + call_depth: usize, + ) -> Result { + self.cursor += 1; + let mut arguments = Vec::new(); + let mut last_was_separator = false; + loop { + self.skip_spaces(); + let token = self.current()?.clone(); + if matches!(token.kind, FormulaTokenKind::RightParen) { + self.cursor += 1; + if last_was_separator { + arguments.push(None); + } + if arguments.len() > MAX_FUNCTION_ARGUMENTS { + return Err(FormulaParseFailure::new( + token.span.start, + format!( + "Function calls accept at most {MAX_FUNCTION_ARGUMENTS} arguments." + ), + )); + } + let child_depth = arguments + .iter() + .filter_map(Option::as_ref) + .map(expression_depth) + .max() + .unwrap_or(0); + return self.node( + SpreadsheetFormulaExpressionKind::FunctionCall { + qualifier, + name, + arguments, + }, + name_span.through(token.span), + child_depth.saturating_add(1), + true, + ); + } + if matches!(token.kind, FormulaTokenKind::End) { + return Err(FormulaParseFailure::new( + token.span.start, + "Function call is not closed with `)`.", + )); + } + if matches!(token.kind, FormulaTokenKind::Comma) { + self.cursor += 1; + arguments.push(None); + last_was_separator = true; + } else { + let argument = self.parse_expression(0, false, call_depth.saturating_add(1))?; + arguments.push(Some(argument.expression)); + last_was_separator = false; + self.skip_spaces(); + let separator = self.current()?.clone(); + match separator.kind { + FormulaTokenKind::Comma => { + self.cursor += 1; + last_was_separator = true; + } + FormulaTokenKind::RightParen => {} + FormulaTokenKind::Semicolon => { + return Err(FormulaParseFailure::new( + separator.span.start, + "SpreadsheetML function arguments use `,`, not `;`.", + )); + } + _ => { + return Err(FormulaParseFailure::new( + separator.span.start, + format!( + "Expected `,` or `)` after a function argument, found {}.", + separator.kind.description() + ), + )); + } + } + } + if arguments.len() >= MAX_FUNCTION_ARGUMENTS && last_was_separator { + return Err(FormulaParseFailure::new( + self.current()?.span.start, + format!("Function calls accept at most {MAX_FUNCTION_ARGUMENTS} arguments."), + )); + } + } + } + + fn parse_array( + &mut self, + open_span: SpreadsheetFormulaSpan, + call_depth: usize, + ) -> Result { + let mut rows = Vec::>::new(); + let mut row = Vec::new(); + let mut last_was_column_separator = false; + loop { + self.skip_spaces(); + let token = self.current()?.clone(); + if matches!(token.kind, FormulaTokenKind::RightBrace) { + if row.is_empty() || last_was_column_separator { + return Err(FormulaParseFailure::new( + token.span.start, + "Array constants cannot contain an empty row.", + )); + } + rows.push(row); + self.cursor += 1; + let width = rows.first().map_or(0, Vec::len); + if rows.iter().any(|value| value.len() != width) { + return Err(FormulaParseFailure::new( + token.span.start, + "Array constant rows must have equal widths.", + )); + } + let child_depth = rows + .iter() + .flatten() + .map(expression_depth) + .max() + .unwrap_or(0); + return self.node( + SpreadsheetFormulaExpressionKind::Array { rows }, + open_span.through(token.span), + child_depth.saturating_add(1), + false, + ); + } + if matches!(token.kind, FormulaTokenKind::End) { + return Err(FormulaParseFailure::new( + token.span.start, + "Array constant is not closed with `}`.", + )); + } + let value = self.parse_expression(0, false, call_depth.saturating_add(1))?; + if !is_array_constant(&value.expression) { + return Err(FormulaParseFailure::new( + value.span().start, + "Array constants may contain only scalar literals.", + )); + } + row.push(value.expression); + last_was_column_separator = false; + self.skip_spaces(); + let separator = self.current()?.clone(); + match separator.kind { + FormulaTokenKind::Comma => { + self.cursor += 1; + last_was_column_separator = true; + } + FormulaTokenKind::Semicolon => { + self.cursor += 1; + if row.is_empty() { + return Err(FormulaParseFailure::new( + separator.span.start, + "Array constants cannot contain an empty row.", + )); + } + rows.push(std::mem::take(&mut row)); + last_was_column_separator = false; + } + FormulaTokenKind::RightBrace => {} + _ => { + return Err(FormulaParseFailure::new( + separator.span.start, + format!( + "Expected `,`, `;`, or `}}` in an array constant, found {}.", + separator.kind.description() + ), + )); + } + } + } + } + + fn binary( + &mut self, + mut left: ParsedExpression, + mut right: ParsedExpression, + operator: SpreadsheetFormulaBinaryOperator, + operator_span: SpreadsheetFormulaSpan, + ) -> Result { + if matches!(operator, SpreadsheetFormulaBinaryOperator::Range) { + left = coerce_range_endpoint(left, operator_span.start)?; + right = coerce_range_endpoint(right, operator_span.start)?; + if let (Some(left_axis), Some(right_axis)) = ( + reference_axis(&left.expression), + reference_axis(&right.expression), + ) { + if left_axis != right_axis { + return Err(FormulaParseFailure::new( + operator_span.start, + "Range endpoints must both be cells, columns, or rows.", + )); + } + } + } else if matches!( + operator, + SpreadsheetFormulaBinaryOperator::Intersection + | SpreadsheetFormulaBinaryOperator::Union + ) && !(left.reference_like && right.reference_like) + { + return Err(FormulaParseFailure::new( + operator_span.start, + "Reference operators require reference expressions on both sides.", + )); + } + let span = left.span().through(right.span()); + let depth = left.depth.max(right.depth).saturating_add(1); + let reference_like = matches!( + operator, + SpreadsheetFormulaBinaryOperator::Range + | SpreadsheetFormulaBinaryOperator::Intersection + | SpreadsheetFormulaBinaryOperator::Union + ); + self.node( + SpreadsheetFormulaExpressionKind::Binary { + operator, + left: Box::new(left.expression), + right: Box::new(right.expression), + }, + span, + depth, + reference_like, + ) + } + + fn node( + &mut self, + kind: SpreadsheetFormulaExpressionKind, + span: SpreadsheetFormulaSpan, + depth: usize, + reference_like: bool, + ) -> Result { + if depth > MAX_SPREADSHEET_FORMULA_DEPTH { + return Err(FormulaParseFailure::new( + span.start, + format!("Formula AST depth exceeds {MAX_SPREADSHEET_FORMULA_DEPTH}."), + )); + } + self.nodes = self.nodes.saturating_add(1); + if self.nodes > MAX_SPREADSHEET_FORMULA_NODES { + return Err(FormulaParseFailure::new( + span.start, + format!("Formula AST contains more than {MAX_SPREADSHEET_FORMULA_NODES} nodes."), + )); + } + Ok(ParsedExpression { + expression: SpreadsheetFormulaExpression { span, kind }, + depth, + reference_like, + }) + } + + fn expect( + &mut self, + expected: FormulaTokenKind, + description: &str, + ) -> Result { + let token = self.current()?.clone(); + if std::mem::discriminant(&token.kind) != std::mem::discriminant(&expected) { + return Err(FormulaParseFailure::new( + token.span.start, + format!( + "Expected {description}, found {}.", + token.kind.description() + ), + )); + } + self.cursor += 1; + Ok(token) + } + + fn current(&self) -> Result<&FormulaToken, FormulaParseFailure> { + self.tokens.get(self.cursor).ok_or_else(|| { + FormulaParseFailure::new( + self.source.len(), + "Formula token stream ended unexpectedly.", + ) + }) + } + + fn skip_spaces(&mut self) { + while self + .tokens + .get(self.cursor) + .is_some_and(|token| matches!(token.kind, FormulaTokenKind::Space)) + { + self.cursor += 1; + } + } + + fn consume_spaces(&mut self) -> Option { + let start = self + .tokens + .get(self.cursor) + .filter(|token| matches!(token.kind, FormulaTokenKind::Space))? + .span; + let mut end = start; + while let Some(token) = self + .tokens + .get(self.cursor) + .filter(|token| matches!(token.kind, FormulaTokenKind::Space)) + { + end = token.span; + self.cursor += 1; + } + Some(start.through(end)) + } + + fn check_call_depth(&self, depth: usize) -> Result<(), FormulaParseFailure> { + if depth > MAX_SPREADSHEET_FORMULA_DEPTH { + return Err(FormulaParseFailure::new( + self.current() + .map_or(self.source.len(), |token| token.span.start), + format!("Formula parse nesting exceeds {MAX_SPREADSHEET_FORMULA_DEPTH}."), + )); + } + Ok(()) + } +} + +fn infix_operator( + token: &FormulaTokenKind, + allow_union: bool, + intersection: bool, +) -> Option<(SpreadsheetFormulaBinaryOperator, u8, u8, bool)> { + if intersection { + return Some(( + SpreadsheetFormulaBinaryOperator::Intersection, + BINDING_INTERSECTION, + BINDING_INTERSECTION + 1, + false, + )); + } + let (operator, binding, right_associative) = match token { + FormulaTokenKind::Colon => ( + SpreadsheetFormulaBinaryOperator::Range, + BINDING_RANGE, + false, + ), + FormulaTokenKind::Comma if allow_union => ( + SpreadsheetFormulaBinaryOperator::Union, + BINDING_UNION, + false, + ), + FormulaTokenKind::Caret => (SpreadsheetFormulaBinaryOperator::Power, BINDING_POWER, true), + FormulaTokenKind::Star => ( + SpreadsheetFormulaBinaryOperator::Multiply, + BINDING_MULTIPLICATIVE, + false, + ), + FormulaTokenKind::Slash => ( + SpreadsheetFormulaBinaryOperator::Divide, + BINDING_MULTIPLICATIVE, + false, + ), + FormulaTokenKind::Plus => ( + SpreadsheetFormulaBinaryOperator::Add, + BINDING_ADDITIVE, + false, + ), + FormulaTokenKind::Minus => ( + SpreadsheetFormulaBinaryOperator::Subtract, + BINDING_ADDITIVE, + false, + ), + FormulaTokenKind::Ampersand => ( + SpreadsheetFormulaBinaryOperator::Concatenate, + BINDING_CONCATENATE, + false, + ), + FormulaTokenKind::Equal => ( + SpreadsheetFormulaBinaryOperator::Equal, + BINDING_COMPARISON, + false, + ), + FormulaTokenKind::NotEqual => ( + SpreadsheetFormulaBinaryOperator::NotEqual, + BINDING_COMPARISON, + false, + ), + FormulaTokenKind::LessThan => ( + SpreadsheetFormulaBinaryOperator::LessThan, + BINDING_COMPARISON, + false, + ), + FormulaTokenKind::LessThanOrEqual => ( + SpreadsheetFormulaBinaryOperator::LessThanOrEqual, + BINDING_COMPARISON, + false, + ), + FormulaTokenKind::GreaterThan => ( + SpreadsheetFormulaBinaryOperator::GreaterThan, + BINDING_COMPARISON, + false, + ), + FormulaTokenKind::GreaterThanOrEqual => ( + SpreadsheetFormulaBinaryOperator::GreaterThanOrEqual, + BINDING_COMPARISON, + false, + ), + _ => return None, + }; + Some(( + operator, + binding, + if right_associative { + binding + } else { + binding + 1 + }, + true, + )) +} + +fn token_starts_reference_expression(token: &FormulaTokenKind) -> bool { + matches!( + token, + FormulaTokenKind::Reference(_) + | FormulaTokenKind::Name { .. } + | FormulaTokenKind::StructuredReference { .. } + | FormulaTokenKind::LeftParen + | FormulaTokenKind::At + ) +} + +fn coerce_range_endpoint( + mut value: ParsedExpression, + position: usize, +) -> Result { + let replacement = match &value.expression.kind { + SpreadsheetFormulaExpressionKind::Name { qualifier, name } => { + if name.bytes().all(|byte| byte.is_ascii_alphabetic()) { + parse_column(name).map(|column| SpreadsheetFormulaReference { + qualifier: qualifier.clone(), + kind: SpreadsheetFormulaReferenceKind::Column { + column, + absolute: false, + }, + }) + } else if name.bytes().all(|byte| byte.is_ascii_digit()) { + parse_row(name).map(|row| SpreadsheetFormulaReference { + qualifier: qualifier.clone(), + kind: SpreadsheetFormulaReferenceKind::Row { + row, + absolute: false, + }, + }) + } else { + None + } + } + SpreadsheetFormulaExpressionKind::Literal(SpreadsheetFormulaLiteral::Number(value)) => { + parse_row(value).map(|row| SpreadsheetFormulaReference { + qualifier: None, + kind: SpreadsheetFormulaReferenceKind::Row { + row, + absolute: false, + }, + }) + } + _ => None, + }; + if let Some(reference) = replacement { + value.expression.kind = SpreadsheetFormulaExpressionKind::Reference(reference); + value.reference_like = true; + } + if !value.reference_like { + return Err(FormulaParseFailure::new( + position, + "Range operator requires reference endpoints.", + )); + } + Ok(value) +} + +fn parse_column(value: &str) -> Option { + if value.is_empty() || value.len() > 3 { + return None; + } + value + .bytes() + .try_fold(0_u32, |column, byte| { + column.checked_mul(26).and_then(|column| { + column.checked_add(u32::from(byte.to_ascii_uppercase().checked_sub(b'A')?) + 1) + }) + }) + .filter(|column| (1..=MAX_COLUMNS).contains(column)) +} + +fn parse_row(value: &str) -> Option { + if value.is_empty() || !value.bytes().all(|byte| byte.is_ascii_digit()) { + return None; + } + value + .parse::() + .ok() + .filter(|row| (1..=MAX_ROWS).contains(row)) +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +enum ReferenceAxisKind { + Cell, + Column, + Row, +} + +fn reference_axis(expression: &SpreadsheetFormulaExpression) -> Option { + let SpreadsheetFormulaExpressionKind::Reference(reference) = &expression.kind else { + return None; + }; + Some(match reference.kind { + SpreadsheetFormulaReferenceKind::Cell { .. } => ReferenceAxisKind::Cell, + SpreadsheetFormulaReferenceKind::Column { .. } => ReferenceAxisKind::Column, + SpreadsheetFormulaReferenceKind::Row { .. } => ReferenceAxisKind::Row, + }) +} + +fn validate_axis_references( + expression: &SpreadsheetFormulaExpression, + range_endpoint: bool, +) -> Result<(), FormulaParseFailure> { + match &expression.kind { + SpreadsheetFormulaExpressionKind::Reference(SpreadsheetFormulaReference { + kind: + SpreadsheetFormulaReferenceKind::Column { .. } + | SpreadsheetFormulaReferenceKind::Row { .. }, + .. + }) if !range_endpoint => Err(FormulaParseFailure::new( + expression.span.start, + "Whole-row and whole-column references must be used in a range.", + )), + SpreadsheetFormulaExpressionKind::Name { + qualifier: Some(_), + name, + } if name.bytes().all(|byte| byte.is_ascii_digit()) => Err(FormulaParseFailure::new( + expression.span.start, + "Qualified row references must include a range operator.", + )), + SpreadsheetFormulaExpressionKind::Unary { operand, .. } + | SpreadsheetFormulaExpressionKind::Postfix { operand, .. } + | SpreadsheetFormulaExpressionKind::Parenthesized(operand) => { + validate_axis_references(operand, false) + } + SpreadsheetFormulaExpressionKind::Binary { + operator, + left, + right, + } => { + let endpoints = matches!(operator, SpreadsheetFormulaBinaryOperator::Range); + validate_axis_references(left, endpoints)?; + validate_axis_references(right, endpoints) + } + SpreadsheetFormulaExpressionKind::FunctionCall { arguments, .. } => { + for argument in arguments.iter().flatten() { + validate_axis_references(argument, false)?; + } + Ok(()) + } + SpreadsheetFormulaExpressionKind::Array { rows } => { + for value in rows.iter().flatten() { + validate_axis_references(value, false)?; + } + Ok(()) + } + _ => Ok(()), + } +} + +fn is_array_constant(expression: &SpreadsheetFormulaExpression) -> bool { + match &expression.kind { + SpreadsheetFormulaExpressionKind::Literal(_) => true, + SpreadsheetFormulaExpressionKind::Unary { + operator: + SpreadsheetFormulaUnaryOperator::Positive | SpreadsheetFormulaUnaryOperator::Negative, + operand, + } => matches!( + &operand.kind, + SpreadsheetFormulaExpressionKind::Literal(SpreadsheetFormulaLiteral::Number(_)) + ), + _ => false, + } +} + +fn expression_depth(expression: &SpreadsheetFormulaExpression) -> usize { + match &expression.kind { + SpreadsheetFormulaExpressionKind::Literal(_) + | SpreadsheetFormulaExpressionKind::Reference(_) + | SpreadsheetFormulaExpressionKind::Name { .. } + | SpreadsheetFormulaExpressionKind::StructuredReference { .. } => 1, + SpreadsheetFormulaExpressionKind::Unary { operand, .. } + | SpreadsheetFormulaExpressionKind::Postfix { operand, .. } + | SpreadsheetFormulaExpressionKind::Parenthesized(operand) => { + expression_depth(operand).saturating_add(1) + } + SpreadsheetFormulaExpressionKind::Binary { left, right, .. } => expression_depth(left) + .max(expression_depth(right)) + .saturating_add(1), + SpreadsheetFormulaExpressionKind::FunctionCall { arguments, .. } => arguments + .iter() + .filter_map(Option::as_ref) + .map(expression_depth) + .max() + .unwrap_or(0) + .saturating_add(1), + SpreadsheetFormulaExpressionKind::Array { rows } => rows + .iter() + .flatten() + .map(expression_depth) + .max() + .unwrap_or(0) + .saturating_add(1), + } +} + +#[cfg(test)] +mod tests; diff --git a/crates/office/src/spreadsheet_formula/parser/tests.rs b/crates/office/src/spreadsheet_formula/parser/tests.rs new file mode 100644 index 00000000..37399915 --- /dev/null +++ b/crates/office/src/spreadsheet_formula/parser/tests.rs @@ -0,0 +1,106 @@ +use super::*; +use crate::spreadsheet_formula::lexer; + +fn parse_formula(source: &str) -> SpreadsheetFormula { + let tokens = lexer::lex(source).unwrap(); + parse(source, tokens).unwrap() +} + +#[test] +fn parser_applies_excel_operator_precedence_and_right_associative_power() { + let formula = parse_formula("-2^3^4%+5*6&\"x\""); + let SpreadsheetFormulaExpressionKind::Binary { + operator: SpreadsheetFormulaBinaryOperator::Concatenate, + left, + right, + } = formula.root.kind + else { + panic!("expected concatenation root"); + }; + assert!(matches!( + right.kind, + SpreadsheetFormulaExpressionKind::Literal(SpreadsheetFormulaLiteral::Text(_)) + )); + let SpreadsheetFormulaExpressionKind::Binary { + operator: SpreadsheetFormulaBinaryOperator::Add, + left: power, + right: multiply, + } = left.kind + else { + panic!("expected additive expression"); + }; + assert!(matches!( + multiply.kind, + SpreadsheetFormulaExpressionKind::Binary { + operator: SpreadsheetFormulaBinaryOperator::Multiply, + .. + } + )); + let SpreadsheetFormulaExpressionKind::Binary { + operator: SpreadsheetFormulaBinaryOperator::Power, + left: negative, + right: nested_power, + } = power.kind + else { + panic!("expected power expression"); + }; + assert!(matches!( + negative.kind, + SpreadsheetFormulaExpressionKind::Unary { + operator: SpreadsheetFormulaUnaryOperator::Negative, + .. + } + )); + assert!(matches!( + nested_power.kind, + SpreadsheetFormulaExpressionKind::Binary { + operator: SpreadsheetFormulaBinaryOperator::Power, + .. + } + )); +} + +#[test] +fn parser_distinguishes_function_arguments_from_reference_unions() { + let formula = parse_formula("SUM(A1:B2,(C1:C2,D1:D2),,TRUE,#N/A)"); + let SpreadsheetFormulaExpressionKind::FunctionCall { arguments, .. } = formula.root.kind else { + panic!("expected function call"); + }; + assert_eq!(arguments.len(), 5); + assert!(arguments[2].is_none()); + let union = arguments[1].as_ref().unwrap(); + let SpreadsheetFormulaExpressionKind::Parenthesized(inner) = &union.kind else { + panic!("expected parenthesized union"); + }; + assert!(matches!( + inner.kind, + SpreadsheetFormulaExpressionKind::Binary { + operator: SpreadsheetFormulaBinaryOperator::Union, + .. + } + )); +} + +#[test] +fn parser_types_cell_column_row_and_sheet_qualified_references() { + let formula = parse_formula("('Q1 Data'!$A$1:B2,[Book.xlsx]Data!C:C) Sheet1!1:$3"); + assert!(matches!( + formula.root.kind, + SpreadsheetFormulaExpressionKind::Binary { + operator: SpreadsheetFormulaBinaryOperator::Intersection, + .. + } + )); +} + +#[test] +fn parser_accepts_rectangular_arrays_and_rejects_invalid_reference_operators() { + assert!(matches!( + parse_formula("{1,-2;TRUE,#N/A}").root.kind, + SpreadsheetFormulaExpressionKind::Array { .. } + )); + for source in ["1 2", "A1::B2", "A:1", "$A+1", "{1,2;3}"] { + let tokens = lexer::lex(source).unwrap(); + assert!(parse(source, tokens).is_err(), "{source}"); + } +} diff --git a/crates/office/src/spreadsheet_formula/registry.rs b/crates/office/src/spreadsheet_formula/registry.rs new file mode 100644 index 00000000..6fb907e4 --- /dev/null +++ b/crates/office/src/spreadsheet_formula/registry.rs @@ -0,0 +1,343 @@ +use std::collections::BTreeMap; + +use serde::{Deserialize, Serialize}; + +/// Native function result family used for registry inspection and validation. +#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)] +#[serde(rename_all = "kebab-case")] +pub enum SpreadsheetFormulaFunctionReturnKind { + Scalar, + ScalarOrArray, + Array, +} + +/// Whether a registered function can change without precedent changes. +#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)] +#[serde(rename_all = "kebab-case")] +pub enum SpreadsheetFormulaFunctionVolatility { + NonVolatile, + Volatile, +} + +/// Closed typed metadata for one natively implemented Spreadsheet function. +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +#[serde(rename_all = "camelCase")] +pub struct SpreadsheetFormulaFunctionDefinition { + pub name: String, + pub minimum_arguments: usize, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub maximum_arguments: Option, + pub return_kind: SpreadsheetFormulaFunctionReturnKind, + pub volatility: SpreadsheetFormulaFunctionVolatility, +} + +/// Deterministic built-in function registry used by native recalculation. +#[derive(Debug, Clone)] +pub struct SpreadsheetFormulaFunctionRegistry { + entries: BTreeMap, +} + +#[derive(Debug, Clone)] +struct FunctionEntry { + definition: SpreadsheetFormulaFunctionDefinition, + function: BuiltinFunction, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub(super) enum BuiltinFunction { + Sum, + Average, + Minimum, + Maximum, + Count, + CountA, + Absolute, + SquareRoot, + Power, + Modulo, + Round, + If, + IfError, + And, + Or, + Not, + Concatenate, + Row, + Column, + Sequence, + Transpose, + Pi, + NotAvailable, +} + +impl Default for SpreadsheetFormulaFunctionRegistry { + fn default() -> Self { + let mut registry = Self { + entries: BTreeMap::new(), + }; + for (name, minimum, maximum, return_kind, function) in [ + ( + "SUM", + 1, + Some(255), + SpreadsheetFormulaFunctionReturnKind::Scalar, + BuiltinFunction::Sum, + ), + ( + "AVERAGE", + 1, + Some(255), + SpreadsheetFormulaFunctionReturnKind::Scalar, + BuiltinFunction::Average, + ), + ( + "MIN", + 1, + Some(255), + SpreadsheetFormulaFunctionReturnKind::Scalar, + BuiltinFunction::Minimum, + ), + ( + "MAX", + 1, + Some(255), + SpreadsheetFormulaFunctionReturnKind::Scalar, + BuiltinFunction::Maximum, + ), + ( + "COUNT", + 1, + Some(255), + SpreadsheetFormulaFunctionReturnKind::Scalar, + BuiltinFunction::Count, + ), + ( + "COUNTA", + 1, + Some(255), + SpreadsheetFormulaFunctionReturnKind::Scalar, + BuiltinFunction::CountA, + ), + ( + "ABS", + 1, + Some(1), + SpreadsheetFormulaFunctionReturnKind::ScalarOrArray, + BuiltinFunction::Absolute, + ), + ( + "SQRT", + 1, + Some(1), + SpreadsheetFormulaFunctionReturnKind::ScalarOrArray, + BuiltinFunction::SquareRoot, + ), + ( + "POWER", + 2, + Some(2), + SpreadsheetFormulaFunctionReturnKind::ScalarOrArray, + BuiltinFunction::Power, + ), + ( + "MOD", + 2, + Some(2), + SpreadsheetFormulaFunctionReturnKind::ScalarOrArray, + BuiltinFunction::Modulo, + ), + ( + "ROUND", + 2, + Some(2), + SpreadsheetFormulaFunctionReturnKind::ScalarOrArray, + BuiltinFunction::Round, + ), + ( + "IF", + 2, + Some(3), + SpreadsheetFormulaFunctionReturnKind::ScalarOrArray, + BuiltinFunction::If, + ), + ( + "IFERROR", + 2, + Some(2), + SpreadsheetFormulaFunctionReturnKind::ScalarOrArray, + BuiltinFunction::IfError, + ), + ( + "AND", + 1, + Some(255), + SpreadsheetFormulaFunctionReturnKind::Scalar, + BuiltinFunction::And, + ), + ( + "OR", + 1, + Some(255), + SpreadsheetFormulaFunctionReturnKind::Scalar, + BuiltinFunction::Or, + ), + ( + "NOT", + 1, + Some(1), + SpreadsheetFormulaFunctionReturnKind::ScalarOrArray, + BuiltinFunction::Not, + ), + ( + "CONCAT", + 1, + Some(255), + SpreadsheetFormulaFunctionReturnKind::Scalar, + BuiltinFunction::Concatenate, + ), + ( + "CONCATENATE", + 1, + Some(255), + SpreadsheetFormulaFunctionReturnKind::Scalar, + BuiltinFunction::Concatenate, + ), + ( + "ROW", + 0, + Some(1), + SpreadsheetFormulaFunctionReturnKind::ScalarOrArray, + BuiltinFunction::Row, + ), + ( + "COLUMN", + 0, + Some(1), + SpreadsheetFormulaFunctionReturnKind::ScalarOrArray, + BuiltinFunction::Column, + ), + ( + "SEQUENCE", + 1, + Some(4), + SpreadsheetFormulaFunctionReturnKind::Array, + BuiltinFunction::Sequence, + ), + ( + "TRANSPOSE", + 1, + Some(1), + SpreadsheetFormulaFunctionReturnKind::Array, + BuiltinFunction::Transpose, + ), + ( + "PI", + 0, + Some(0), + SpreadsheetFormulaFunctionReturnKind::Scalar, + BuiltinFunction::Pi, + ), + ( + "NA", + 0, + Some(0), + SpreadsheetFormulaFunctionReturnKind::Scalar, + BuiltinFunction::NotAvailable, + ), + ] { + registry.insert( + name, + minimum, + maximum, + return_kind, + SpreadsheetFormulaFunctionVolatility::NonVolatile, + function, + ); + } + registry + } +} + +impl SpreadsheetFormulaFunctionRegistry { + pub fn definitions( + &self, + ) -> impl ExactSizeIterator { + self.entries.values().map(|entry| &entry.definition) + } + + pub fn get(&self, name: &str) -> Option<&SpreadsheetFormulaFunctionDefinition> { + self.resolve(name).map(|entry| &entry.definition) + } + + pub fn contains(&self, name: &str) -> bool { + self.resolve(name).is_some() + } + + pub(super) fn function(&self, name: &str) -> Option { + self.resolve(name).map(|entry| entry.function) + } + + fn resolve(&self, name: &str) -> Option<&FunctionEntry> { + self.entries.get(&normalize_function_name(name)) + } + + fn insert( + &mut self, + name: &str, + minimum_arguments: usize, + maximum_arguments: Option, + return_kind: SpreadsheetFormulaFunctionReturnKind, + volatility: SpreadsheetFormulaFunctionVolatility, + function: BuiltinFunction, + ) { + self.entries.insert( + name.to_string(), + FunctionEntry { + definition: SpreadsheetFormulaFunctionDefinition { + name: name.to_string(), + minimum_arguments, + maximum_arguments, + return_kind, + volatility, + }, + function, + }, + ); + } +} + +pub(super) fn normalize_function_name(name: &str) -> String { + let mut normalized = name.to_ascii_uppercase(); + loop { + let stripped = ["_XLFN.", "_XLWS."] + .into_iter() + .find_map(|prefix| normalized.strip_prefix(prefix).map(ToOwned::to_owned)); + let Some(stripped) = stripped else { + return normalized; + }; + normalized = stripped; + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn built_in_registry_is_typed_bounded_and_namespace_aware() { + fn assert_send_sync() {} + assert_send_sync::(); + + let registry = SpreadsheetFormulaFunctionRegistry::default(); + assert!(registry.contains("sum")); + assert!(registry.contains("_xlfn._xlws.SEQUENCE")); + assert!(!registry.contains("SHELL")); + let sequence = registry.get("SEQUENCE").unwrap(); + assert_eq!(sequence.minimum_arguments, 1); + assert_eq!(sequence.maximum_arguments, Some(4)); + assert_eq!( + sequence.return_kind, + SpreadsheetFormulaFunctionReturnKind::Array + ); + } +} diff --git a/crates/office/src/spreadsheet_formula/structured_reference.rs b/crates/office/src/spreadsheet_formula/structured_reference.rs new file mode 100644 index 00000000..f5203afd --- /dev/null +++ b/crates/office/src/spreadsheet_formula/structured_reference.rs @@ -0,0 +1,516 @@ +mod parser; +mod rewrite; + +use std::collections::BTreeMap; + +use a3s_use_core::{UseError, UseResult}; + +use crate::discovery::office_error; +use crate::semantic::{DocumentNode, OfficeNodeType}; +use crate::spreadsheet_reference::{CellRange, CellReference}; + +use super::SpreadsheetFormulaQualifier; +use parser::{parse_reference, ParsedStructuredReference, StructuredRowSelection}; +pub(crate) use rewrite::{ + LocalStructuredReferenceContext, StructuredReferenceRewritePlan, + StructuredReferenceRewriteResult, +}; + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub(crate) enum StructuredReferenceErrorKind { + ExternalWorkbook, + Unsupported, + MissingTable, + MissingColumn, +} + +#[derive(Debug, Clone)] +pub(crate) struct StructuredReferenceError { + pub(crate) kind: StructuredReferenceErrorKind, + pub(crate) message: String, +} + +#[derive(Debug, Clone, Copy)] +pub(crate) struct ResolvedStructuredReference { + pub(crate) sheet: usize, + pub(crate) start_column: u32, + pub(crate) start_row: u32, + pub(crate) end_column: u32, + pub(crate) end_row: u32, +} + +#[derive(Debug, Clone)] +struct FormulaTableDefinition { + path: String, + name: String, + sheet: usize, + range: CellRange, + header_row: bool, + totals_row: bool, + columns: Vec, +} + +#[derive(Debug, Clone)] +pub(crate) struct FormulaTableCatalog { + sheet_names: Vec, + definitions: Vec, + by_name: BTreeMap, + by_sheet: Vec>, +} + +impl FormulaTableCatalog { + pub(crate) fn collect(root: &DocumentNode, sheet_names: &[String]) -> UseResult { + let mut definitions = Vec::::new(); + let mut by_name = BTreeMap::::new(); + let mut by_sheet = vec![Vec::::new(); sheet_names.len()]; + for (sheet, sheet_name) in sheet_names.iter().enumerate() { + let worksheet = root + .children + .iter() + .find(|node| { + node.node_type == OfficeNodeType::Worksheet + && node + .path + .strip_prefix('/') + .is_some_and(|path| path.eq_ignore_ascii_case(sheet_name)) + }) + .ok_or_else(|| table_catalog_error(format!("Missing worksheet '{sheet_name}'.")))?; + for table in worksheet + .children + .iter() + .filter(|node| node.node_type == OfficeNodeType::Table) + { + let name = required_table_format(table, "name")?; + let display_name = required_table_format(table, "displayName")?; + let range_text = required_table_format(table, "ref")?; + let range = CellRange::parse(range_text).map_err(|error| { + table_catalog_error(format!( + "Spreadsheet table '{}' has invalid range '{range_text}': {error}", + table.path + )) + })?; + let header_row = table_boolean(table, "headerRow")?; + let totals_row = table_boolean(table, "totalsRow")?; + let columns = table + .children + .iter() + .filter(|node| node.node_type == OfficeNodeType::TableColumn) + .map(|column| { + column.format.get("name").cloned().ok_or_else(|| { + table_catalog_error(format!( + "Spreadsheet table column '{}' has no name.", + column.path + )) + }) + }) + .collect::>>()?; + let width = usize::try_from(range.end.column - range.start.column + 1) + .map_err(|_| table_catalog_error("Spreadsheet table width is invalid."))?; + if columns.len() != width { + return Err(table_catalog_error(format!( + "Spreadsheet table '{}' has {} columns for range '{}'.", + table.path, + columns.len(), + range.a1() + ))); + } + let definition = FormulaTableDefinition { + path: table.path.clone(), + name: name.to_string(), + sheet, + range, + header_row, + totals_row, + columns, + }; + let aliases = [name, display_name] + .into_iter() + .map(|alias| (alias.to_string(), alias.to_lowercase())) + .collect::>(); + for (alias, normalized) in &aliases { + if let Some(existing) = by_name.get(normalized) { + let existing = definitions.get(*existing).ok_or_else(|| { + table_catalog_error("Spreadsheet table name index is invalid.") + })?; + if existing.path != definition.path { + return Err(table_catalog_error(format!( + "Spreadsheet table formula name '{alias}' is ambiguous." + ))); + } + } + } + let definition_index = definitions.len(); + definitions.push(definition); + by_sheet + .get_mut(sheet) + .ok_or_else(|| { + table_catalog_error("Spreadsheet table worksheet index is invalid.") + })? + .push(definition_index); + for (_, normalized) in aliases { + by_name.entry(normalized).or_insert(definition_index); + } + } + } + Ok(Self { + sheet_names: sheet_names.to_vec(), + definitions, + by_name, + by_sheet, + }) + } + + pub(crate) fn resolve( + &self, + qualifier: Option<&SpreadsheetFormulaQualifier>, + reference: &str, + current_sheet: usize, + current_column: u32, + current_row: u32, + ) -> Result, StructuredReferenceError> { + let parsed = parse_reference(reference)?; + let qualifier_sheet = self.resolve_qualifier(qualifier)?; + let table = self.resolve_table( + &parsed, + reference, + current_sheet, + current_column, + current_row, + )?; + if let Some(sheet) = qualifier_sheet { + if sheet != table.sheet { + return Err(structured_error( + StructuredReferenceErrorKind::MissingTable, + format!( + "Spreadsheet table '{}' is not on worksheet '{}'.", + table.name, + qualifier + .map(|value| value.worksheet.as_str()) + .unwrap_or_default() + ), + )); + } + } + if parsed.rows.current && current_sheet != table.sheet { + return Err(structured_error( + StructuredReferenceErrorKind::Unsupported, + format!( + "Structured reference #This Row requires the current formula cell to be on table '{}'.", + table.name + ), + )); + } + let (start_column, end_column) = resolve_columns(table, &parsed)?; + resolve_rows(table, parsed.rows, current_row)? + .into_iter() + .map(|(start_row, end_row)| { + Ok(ResolvedStructuredReference { + sheet: table.sheet, + start_column, + start_row, + end_column, + end_row, + }) + }) + .collect() + } + + fn resolve_qualifier( + &self, + qualifier: Option<&SpreadsheetFormulaQualifier>, + ) -> Result, StructuredReferenceError> { + let Some(qualifier) = qualifier else { + return Ok(None); + }; + if qualifier.is_external() { + return Err(structured_error( + StructuredReferenceErrorKind::ExternalWorkbook, + "Structured references to external workbooks are not supported.", + )); + } + if qualifier.is_three_dimensional() { + return Err(structured_error( + StructuredReferenceErrorKind::Unsupported, + "Three-dimensional structured-reference qualifiers are not supported.", + )); + } + self.sheet_names + .iter() + .position(|name| name.eq_ignore_ascii_case(&qualifier.worksheet)) + .map(Some) + .ok_or_else(|| { + structured_error( + StructuredReferenceErrorKind::MissingTable, + format!( + "Structured-reference worksheet '{}' does not exist.", + qualifier.worksheet + ), + ) + }) + } + + fn resolve_table( + &self, + parsed: &ParsedStructuredReference, + reference: &str, + current_sheet: usize, + current_column: u32, + current_row: u32, + ) -> Result<&FormulaTableDefinition, StructuredReferenceError> { + if let Some(table_name) = &parsed.table_name { + let index = self + .by_name + .get(&table_name.to_lowercase()) + .ok_or_else(|| { + structured_error( + StructuredReferenceErrorKind::MissingTable, + format!("Spreadsheet table '{table_name}' does not exist."), + ) + })?; + return self.definitions.get(*index).ok_or_else(|| { + structured_error( + StructuredReferenceErrorKind::Unsupported, + "Structured-reference table index is invalid.", + ) + }); + } + + let current = CellReference { + column: current_column, + row: current_row, + }; + let mut matching = self + .by_sheet + .get(current_sheet) + .into_iter() + .flatten() + .filter_map(|index| self.definitions.get(*index)) + .filter(|table| table.range.contains(current)); + let Some(table) = matching.next() else { + return Err(structured_error( + StructuredReferenceErrorKind::MissingTable, + format!( + "Table-local structured reference '{reference}' requires its formula cell to be inside a Spreadsheet table." + ), + )); + }; + if matching.next().is_some() { + return Err(structured_error( + StructuredReferenceErrorKind::Unsupported, + format!( + "Table-local structured reference '{reference}' is ambiguous at the current formula cell." + ), + )); + } + Ok(table) + } +} + +fn resolve_columns( + table: &FormulaTableDefinition, + parsed: &ParsedStructuredReference, +) -> Result<(u32, u32), StructuredReferenceError> { + let (first, last) = match (&parsed.first_column, &parsed.last_column) { + (None, None) => (0, table.columns.len().saturating_sub(1)), + (Some(first), Some(last)) => { + let first_index = table + .columns + .iter() + .position(|name| name.eq_ignore_ascii_case(first)) + .ok_or_else(|| missing_column(&table.name, first))?; + let last_index = table + .columns + .iter() + .position(|name| name.eq_ignore_ascii_case(last)) + .ok_or_else(|| missing_column(&table.name, last))?; + if first_index > last_index { + return Err(structured_error( + StructuredReferenceErrorKind::Unsupported, + format!("Structured-reference column range '{first}:{last}' is reversed."), + )); + } + (first_index, last_index) + } + _ => { + return Err(structured_error( + StructuredReferenceErrorKind::Unsupported, + "Structured-reference column selection is incomplete.", + )) + } + }; + let first = u32::try_from(first).map_err(|_| invalid_column_index())?; + let last = u32::try_from(last).map_err(|_| invalid_column_index())?; + let start_column = table + .range + .start + .column + .checked_add(first) + .ok_or_else(invalid_column_index)?; + let end_column = table + .range + .start + .column + .checked_add(last) + .ok_or_else(invalid_column_index)?; + if start_column > table.range.end.column || end_column > table.range.end.column { + return Err(invalid_column_index()); + } + Ok((start_column, end_column)) +} + +fn resolve_rows( + table: &FormulaTableDefinition, + rows: StructuredRowSelection, + current_row: u32, +) -> Result, StructuredReferenceError> { + let mut selected = Vec::<(u32, u32)>::new(); + if rows.all { + selected.push((table.range.start.row, table.range.end.row)); + } + if rows.headers { + if !table.header_row { + return Err(missing_table_rows(table, "#Headers")); + } + selected.push((table.range.start.row, table.range.start.row)); + } + if rows.data { + selected.push(data_rows(table)?); + } + if rows.totals { + if !table.totals_row { + return Err(missing_table_rows(table, "#Totals")); + } + selected.push((table.range.end.row, table.range.end.row)); + } + if rows.current { + let (start, end) = data_rows(table)?; + if !(start..=end).contains(¤t_row) { + return Err(structured_error( + StructuredReferenceErrorKind::Unsupported, + format!( + "Structured reference #This Row requires the current formula row to be inside table '{}'.", + table.name + ), + )); + } + selected.push((current_row, current_row)); + } + selected.sort_unstable(); + let mut merged = Vec::<(u32, u32)>::with_capacity(selected.len()); + for (start, end) in selected { + if let Some((_, previous_end)) = merged.last_mut() { + if start <= previous_end.saturating_add(1) { + *previous_end = (*previous_end).max(end); + continue; + } + } + merged.push((start, end)); + } + if merged.is_empty() { + return Err(structured_error( + StructuredReferenceErrorKind::Unsupported, + "Structured reference selects no table rows.", + )); + } + Ok(merged) +} + +fn data_rows(table: &FormulaTableDefinition) -> Result<(u32, u32), StructuredReferenceError> { + let start = table + .range + .start + .row + .checked_add(u32::from(table.header_row)) + .ok_or_else(|| { + structured_error( + StructuredReferenceErrorKind::Unsupported, + "Structured-reference data range exceeds worksheet limits.", + ) + })?; + let end = table + .range + .end + .row + .checked_sub(u32::from(table.totals_row)) + .ok_or_else(|| { + structured_error( + StructuredReferenceErrorKind::Unsupported, + "Structured-reference data range is invalid.", + ) + })?; + if start > end { + return Err(structured_error( + StructuredReferenceErrorKind::Unsupported, + format!("Spreadsheet table '{}' has no data rows.", table.name), + )); + } + Ok((start, end)) +} + +fn missing_table_rows(table: &FormulaTableDefinition, item: &str) -> StructuredReferenceError { + structured_error( + StructuredReferenceErrorKind::Unsupported, + format!( + "Structured reference {item} requires table '{}' to contain that row.", + table.name + ), + ) +} + +fn invalid_column_index() -> StructuredReferenceError { + structured_error( + StructuredReferenceErrorKind::Unsupported, + "Structured-reference column index is invalid.", + ) +} + +fn required_table_format<'a>(table: &'a DocumentNode, key: &str) -> UseResult<&'a str> { + table.format.get(key).map(String::as_str).ok_or_else(|| { + table_catalog_error(format!( + "Spreadsheet table '{}' has no '{key}' property.", + table.path + )) + }) +} + +fn table_boolean(table: &DocumentNode, key: &str) -> UseResult { + match required_table_format(table, key)? { + "true" => Ok(true), + "false" => Ok(false), + value => Err(table_catalog_error(format!( + "Spreadsheet table '{}' has invalid boolean '{key}={value}'.", + table.path + ))), + } +} + +fn missing_column(table: &str, column: &str) -> StructuredReferenceError { + structured_error( + StructuredReferenceErrorKind::MissingColumn, + format!("Spreadsheet table '{table}' has no column '{column}'."), + ) +} + +fn invalid_reference(reference: &str) -> StructuredReferenceError { + structured_error( + StructuredReferenceErrorKind::Unsupported, + format!("Structured reference '{reference}' is not in a supported canonical form."), + ) +} + +fn structured_error( + kind: StructuredReferenceErrorKind, + message: impl Into, +) -> StructuredReferenceError { + StructuredReferenceError { + kind, + message: message.into(), + } +} + +fn table_catalog_error(message: impl Into) -> UseError { + office_error( + "use.office.spreadsheet_formula_table_catalog_invalid", + message, + ) +} diff --git a/crates/office/src/spreadsheet_formula/structured_reference/parser.rs b/crates/office/src/spreadsheet_formula/structured_reference/parser.rs new file mode 100644 index 00000000..d85dffef --- /dev/null +++ b/crates/office/src/spreadsheet_formula/structured_reference/parser.rs @@ -0,0 +1,270 @@ +use super::{ + invalid_reference, structured_error, StructuredReferenceError, StructuredReferenceErrorKind, +}; + +#[derive(Debug, Clone, Copy, Default)] +pub(super) struct StructuredRowSelection { + pub(super) all: bool, + pub(super) headers: bool, + pub(super) data: bool, + pub(super) totals: bool, + pub(super) current: bool, +} + +#[derive(Debug, Clone)] +pub(super) struct ParsedStructuredReference { + pub(super) table_name: Option, + pub(super) first_column: Option, + pub(super) last_column: Option, + pub(super) rows: StructuredRowSelection, +} + +pub(super) fn parse_reference( + reference: &str, +) -> Result { + let Some(open) = reference.find('[') else { + return Err(invalid_reference(reference)); + }; + let table_name = (!reference[..open].is_empty()).then(|| reference[..open].to_string()); + let content = outer_group(&reference[open..]).ok_or_else(|| invalid_reference(reference))?; + let mut rows = StructuredRowSelection::default(); + let (first_column, last_column) = if let Some(current) = content.strip_prefix('@') { + rows.current = true; + let column = parse_current_column(current, reference)?; + (Some(column.clone()), Some(column)) + } else if content.starts_with('[') { + parse_nested_selection(content, reference, &mut rows)? + } else if let Some(item) = table_item(content) { + apply_table_item(&mut rows, item, reference)?; + (None, None) + } else { + let column = parse_plain_column(content, reference)?; + (Some(column.clone()), Some(column)) + }; + if !rows.all && !rows.headers && !rows.data && !rows.totals && !rows.current { + rows.data = true; + } + Ok(ParsedStructuredReference { + table_name, + first_column, + last_column, + rows, + }) +} + +fn parse_current_column(value: &str, reference: &str) -> Result { + if value.starts_with('[') { + let (atom, consumed) = bracket_atom(value).ok_or_else(|| invalid_reference(reference))?; + if consumed != value.len() { + return Err(invalid_reference(reference)); + } + let column = decode_atom(atom)?; + if column.is_empty() { + return Err(invalid_reference(reference)); + } + return Ok(column); + } + parse_plain_column(value, reference) +} + +fn parse_nested_selection( + content: &str, + reference: &str, + rows: &mut StructuredRowSelection, +) -> Result<(Option, Option), StructuredReferenceError> { + let (atoms, separators) = nested_components(content, reference)?; + let mut columns = Vec::<(usize, String)>::new(); + for (index, atom) in atoms.iter().enumerate() { + if let Some(item) = table_item(atom) { + apply_table_item(rows, item, reference)?; + } else { + let column = decode_atom(atom)?; + if column.is_empty() { + return Err(invalid_reference(reference)); + } + columns.push((index, column)); + } + } + match columns.as_slice() { + [] => { + if separators.contains(&':') { + return Err(invalid_reference(reference)); + } + Ok((None, None)) + } + [(_, column)] => { + if separators.contains(&':') { + return Err(invalid_reference(reference)); + } + Ok((Some(column.clone()), Some(column.clone()))) + } + [(first_index, first), (last_index, last)] + if *last_index == first_index.saturating_add(1) + && separators.get(*first_index) == Some(&':') + && separators + .iter() + .enumerate() + .all(|(index, separator)| index == *first_index || *separator == ',') => + { + Ok((Some(first.clone()), Some(last.clone()))) + } + _ => Err(structured_error( + StructuredReferenceErrorKind::Unsupported, + "Disjoint structured-reference columns are not supported.", + )), + } +} + +fn nested_components<'a>( + content: &'a str, + reference: &str, +) -> Result<(Vec<&'a str>, Vec), StructuredReferenceError> { + let mut atoms = Vec::new(); + let mut separators = Vec::new(); + let mut cursor = 0_usize; + loop { + let (atom, consumed) = + bracket_atom(&content[cursor..]).ok_or_else(|| invalid_reference(reference))?; + atoms.push(atom); + cursor = cursor + .checked_add(consumed) + .ok_or_else(|| invalid_reference(reference))?; + if cursor == content.len() { + break; + } + let separator = content[cursor..] + .chars() + .next() + .ok_or_else(|| invalid_reference(reference))?; + if !matches!(separator, ',' | ':') { + return Err(invalid_reference(reference)); + } + separators.push(separator); + cursor = cursor + .checked_add(separator.len_utf8()) + .ok_or_else(|| invalid_reference(reference))?; + if cursor >= content.len() { + return Err(invalid_reference(reference)); + } + } + Ok((atoms, separators)) +} + +fn parse_plain_column(value: &str, reference: &str) -> Result { + if value.is_empty() || value.contains(['[', ']', ',', ':']) { + return Err(invalid_reference(reference)); + } + let column = decode_atom(value)?; + if column.is_empty() { + return Err(invalid_reference(reference)); + } + Ok(column) +} + +#[derive(Debug, Clone, Copy)] +enum TableItem { + All, + Headers, + Data, + Totals, + Current, +} + +fn table_item(value: &str) -> Option { + [ + ("#all", TableItem::All), + ("#headers", TableItem::Headers), + ("#data", TableItem::Data), + ("#totals", TableItem::Totals), + ("#this row", TableItem::Current), + ] + .into_iter() + .find_map(|(name, item)| value.eq_ignore_ascii_case(name).then_some(item)) +} + +fn apply_table_item( + rows: &mut StructuredRowSelection, + item: TableItem, + reference: &str, +) -> Result<(), StructuredReferenceError> { + match item { + TableItem::All => rows.all = true, + TableItem::Headers => rows.headers = true, + TableItem::Data => rows.data = true, + TableItem::Totals => rows.totals = true, + TableItem::Current => rows.current = true, + } + if rows.current && (rows.all || rows.headers || rows.data || rows.totals) { + return Err(structured_error( + StructuredReferenceErrorKind::Unsupported, + format!( + "Structured reference '{reference}' cannot combine #This Row with another item." + ), + )); + } + Ok(()) +} + +fn outer_group(value: &str) -> Option<&str> { + if !value.starts_with('[') { + return None; + } + let end = matching_bracket(value, 0)?; + (end == value.len()).then_some(&value[1..value.len() - 1]) +} + +fn bracket_atom(value: &str) -> Option<(&str, usize)> { + if !value.starts_with('[') { + return None; + } + let end = matching_bracket(value, 0)?; + Some((&value[1..end - 1], end)) +} + +fn matching_bracket(value: &str, start: usize) -> Option { + let mut depth = 0_usize; + let mut cursor = start; + while cursor < value.len() { + let character = value[cursor..].chars().next()?; + cursor += character.len_utf8(); + if character == '\'' { + let escaped = value[cursor..].chars().next()?; + cursor += escaped.len_utf8(); + continue; + } + match character { + '[' => depth = depth.checked_add(1)?, + ']' => { + depth = depth.checked_sub(1)?; + if depth == 0 { + return Some(cursor); + } + } + _ => {} + } + } + None +} + +fn decode_atom(value: &str) -> Result { + let mut output = String::with_capacity(value.len()); + let mut cursor = 0_usize; + while cursor < value.len() { + let character = value[cursor..] + .chars() + .next() + .ok_or_else(|| invalid_reference(value))?; + cursor += character.len_utf8(); + if character == '\'' { + let escaped = value[cursor..] + .chars() + .next() + .ok_or_else(|| invalid_reference(value))?; + cursor += escaped.len_utf8(); + output.push(escaped); + } else { + output.push(character); + } + } + Ok(output) +} diff --git a/crates/office/src/spreadsheet_formula/structured_reference/rewrite.rs b/crates/office/src/spreadsheet_formula/structured_reference/rewrite.rs new file mode 100644 index 00000000..ed74bcf2 --- /dev/null +++ b/crates/office/src/spreadsheet_formula/structured_reference/rewrite.rs @@ -0,0 +1,538 @@ +use std::collections::BTreeMap; +use std::ops::Range; + +use a3s_use_core::{UseError, UseResult}; + +use crate::discovery::office_error; + +use super::super::lexer::{self, FormulaTokenKind}; +use super::super::parse_error; + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub(crate) enum LocalStructuredReferenceContext { + Applies, + DoesNotApply, + Unknown, +} + +#[derive(Debug, Clone)] +pub(crate) struct StructuredReferenceRewritePlan { + table_name: String, + table_sheet: String, + aliases: BTreeMap, + columns: BTreeMap>, + geometry_changed: bool, + removal: bool, +} + +#[derive(Debug, Clone)] +pub(crate) struct StructuredReferenceRewriteResult { + pub(crate) formula: String, + pub(crate) matched: bool, +} + +impl StructuredReferenceRewritePlan { + pub(crate) fn rename( + table_name: impl Into, + table_sheet: impl Into, + aliases: BTreeMap, + columns: BTreeMap>, + geometry_changed: bool, + ) -> Self { + Self { + table_name: table_name.into(), + table_sheet: table_sheet.into(), + aliases, + columns, + geometry_changed, + removal: false, + } + } + + pub(crate) fn removal( + table_name: impl Into, + table_sheet: impl Into, + aliases: impl IntoIterator, + ) -> Self { + Self { + table_name: table_name.into(), + table_sheet: table_sheet.into(), + aliases: aliases + .into_iter() + .map(|alias| (alias.to_lowercase(), alias)) + .collect(), + columns: BTreeMap::new(), + geometry_changed: false, + removal: true, + } + } + + pub(crate) const fn geometry_changed(&self) -> bool { + self.geometry_changed + } + + pub(crate) fn rewrite( + &self, + formula: &str, + local_context: LocalStructuredReferenceContext, + ) -> UseResult { + if !formula.contains('[') { + return Ok(StructuredReferenceRewriteResult { + formula: formula.to_string(), + matched: false, + }); + } + let body_offset = usize::from(formula.starts_with('=')); + let body = &formula[body_offset..]; + let tokens = lexer::lex(body).map_err(|failure| parse_error(body, failure))?; + let mut replacements = Vec::<(Range, String)>::new(); + let mut matched = false; + for token in tokens { + let FormulaTokenKind::StructuredReference { + qualifier, + reference, + } = token.kind + else { + continue; + }; + let Some(open) = reference.find('[') else { + continue; + }; + let reference_start = body_offset + .checked_add(token.span.end.saturating_sub(reference.len())) + .ok_or_else(rewrite_limit_error)?; + let table_name = &reference[..open]; + if table_name.is_empty() { + let local_matched = self.rewrite_local( + &reference, + reference_start, + local_context, + &mut replacements, + )?; + matched |= local_matched; + continue; + } + let Some(replacement_name) = self.aliases.get(&table_name.to_lowercase()) else { + continue; + }; + if qualifier + .as_ref() + .is_some_and(super::super::SpreadsheetFormulaQualifier::is_external) + { + continue; + } + if let Some(qualifier) = qualifier.as_ref() { + if qualifier.is_three_dimensional() { + return Err(rewrite_unsupported( + &self.table_name, + "Three-dimensional structured references cannot be rewritten safely.", + )); + } + if !qualifier.worksheet.eq_ignore_ascii_case(&self.table_sheet) { + continue; + } + } + matched = true; + if self.removal { + return Err(table_referenced(&self.table_name, &reference)); + } + if replacement_name != table_name { + replacements.push(( + reference_start..reference_start + open, + replacement_name.clone(), + )); + } + self.rewrite_columns(&reference, reference_start, &mut replacements)?; + } + Ok(StructuredReferenceRewriteResult { + formula: apply_replacements(formula, replacements)?, + matched, + }) + } + + fn rewrite_local( + &self, + reference: &str, + reference_start: usize, + context: LocalStructuredReferenceContext, + replacements: &mut Vec<(Range, String)>, + ) -> UseResult { + if matches!(context, LocalStructuredReferenceContext::DoesNotApply) { + return Ok(false); + } + if self.removal { + return Err(table_referenced(&self.table_name, reference)); + } + if self.geometry_changed { + return Err(rewrite_unsupported( + &self.table_name, + "Table-local structured references cannot be retained across table geometry or structural-row changes.", + )); + } + let atoms = column_atoms(reference)?; + let affected = atoms + .iter() + .any(|atom| self.columns.contains_key(&atom.name.to_lowercase())); + if !affected { + return Ok(matches!(context, LocalStructuredReferenceContext::Applies)); + } + if matches!(context, LocalStructuredReferenceContext::Unknown) { + return Err(rewrite_unsupported( + &self.table_name, + "A table-local structured reference has no provable ListObject context.", + )); + } + self.rewrite_column_atoms(atoms, reference_start, replacements)?; + Ok(true) + } + + fn rewrite_columns( + &self, + reference: &str, + reference_start: usize, + replacements: &mut Vec<(Range, String)>, + ) -> UseResult<()> { + if self.columns.is_empty() { + return Ok(()); + } + self.rewrite_column_atoms(column_atoms(reference)?, reference_start, replacements) + } + + fn rewrite_column_atoms( + &self, + atoms: Vec, + reference_start: usize, + replacements: &mut Vec<(Range, String)>, + ) -> UseResult<()> { + for atom in atoms { + let Some(replacement) = self.columns.get(&atom.name.to_lowercase()) else { + continue; + }; + let Some(replacement) = replacement else { + return Err(rewrite_unsupported( + &self.table_name, + format!( + "Structured reference column '{}' would be removed from the table.", + atom.name + ), + )); + }; + if replacement == &atom.name { + continue; + } + let encoded = encode_column(replacement, atom.plain, atom.current)?; + replacements.push(( + reference_start + atom.range.start..reference_start + atom.range.end, + encoded, + )); + } + Ok(()) + } +} + +#[derive(Debug, Clone)] +struct ColumnAtom { + name: String, + range: Range, + plain: bool, + current: bool, +} + +fn column_atoms(reference: &str) -> UseResult> { + let Some(mut cursor) = reference.find('[') else { + return Ok(Vec::new()); + }; + let mut atoms = Vec::new(); + while cursor < reference.len() { + if !reference[cursor..].starts_with('[') { + return Err(rewrite_unsupported( + reference, + "Structured-reference bracket groups are not canonical.", + )); + } + cursor = collect_group(reference, cursor, 1, &mut atoms)?; + } + Ok(atoms) +} + +fn collect_group( + reference: &str, + start: usize, + depth: usize, + atoms: &mut Vec, +) -> UseResult { + let mut cursor = start.checked_add(1).ok_or_else(rewrite_limit_error)?; + let content_start = cursor; + let mut has_child = false; + while cursor < reference.len() { + let character = reference[cursor..] + .chars() + .next() + .ok_or_else(rewrite_limit_error)?; + if character == '\'' { + cursor = cursor + .checked_add(character.len_utf8()) + .ok_or_else(rewrite_limit_error)?; + let escaped = reference[cursor..].chars().next().ok_or_else(|| { + rewrite_unsupported( + reference, + "Structured-reference escape has no following character.", + ) + })?; + cursor = cursor + .checked_add(escaped.len_utf8()) + .ok_or_else(rewrite_limit_error)?; + continue; + } + match character { + '[' => { + has_child = true; + cursor = collect_group(reference, cursor, depth.saturating_add(1), atoms)?; + } + ']' => { + if !has_child { + push_column_atom(reference, content_start..cursor, depth, atoms)?; + } + return cursor + .checked_add(character.len_utf8()) + .ok_or_else(rewrite_limit_error); + } + _ => { + cursor = cursor + .checked_add(character.len_utf8()) + .ok_or_else(rewrite_limit_error)?; + } + } + } + Err(rewrite_unsupported( + reference, + "Structured-reference bracket is not closed.", + )) +} + +fn push_column_atom( + reference: &str, + mut range: Range, + depth: usize, + atoms: &mut Vec, +) -> UseResult<()> { + let raw = reference + .get(range.clone()) + .ok_or_else(rewrite_limit_error)?; + if table_item(raw) { + return Ok(()); + } + let current = depth == 1 && raw.starts_with('@'); + if current { + range.start = range.start.checked_add(1).ok_or_else(rewrite_limit_error)?; + } + let raw = reference + .get(range.clone()) + .ok_or_else(rewrite_limit_error)?; + if raw.is_empty() { + return Err(rewrite_unsupported( + reference, + "Structured-reference column is empty.", + )); + } + atoms.push(ColumnAtom { + name: decode_atom(raw)?, + range, + plain: depth == 1, + current, + }); + Ok(()) +} + +fn decode_atom(raw: &str) -> UseResult { + let mut output = String::with_capacity(raw.len()); + let mut cursor = 0_usize; + while cursor < raw.len() { + let character = raw[cursor..] + .chars() + .next() + .ok_or_else(rewrite_limit_error)?; + cursor = cursor + .checked_add(character.len_utf8()) + .ok_or_else(rewrite_limit_error)?; + if character == '\'' { + let escaped = raw[cursor..].chars().next().ok_or_else(|| { + rewrite_unsupported( + raw, + "Structured-reference escape has no following character.", + ) + })?; + cursor = cursor + .checked_add(escaped.len_utf8()) + .ok_or_else(rewrite_limit_error)?; + output.push(escaped); + } else { + output.push(character); + } + } + Ok(output) +} + +fn encode_column(value: &str, plain: bool, current: bool) -> UseResult { + if plain && value.contains(['[', ']', ',', ':']) { + return Err(rewrite_unsupported( + value, + "The replacement column requires a nested structured-reference form.", + )); + } + let mut output = String::with_capacity(value.len()); + for (index, character) in value.chars().enumerate() { + let escape_leading = index == 0 + && ((plain && !current && matches!(character, '#' | '@')) + || (!plain && character == '#')); + if escape_leading || matches!(character, '\'' | '[' | ']') { + output.push('\''); + } + output.push(character); + } + Ok(output) +} + +fn table_item(value: &str) -> bool { + ["#all", "#headers", "#data", "#totals", "#this row"] + .into_iter() + .any(|item| value.eq_ignore_ascii_case(item)) +} + +fn apply_replacements( + formula: &str, + mut replacements: Vec<(Range, String)>, +) -> UseResult { + if replacements.is_empty() { + return Ok(formula.to_string()); + } + replacements.sort_by_key(|(range, _)| (range.start, range.end)); + if replacements + .windows(2) + .any(|pair| pair[0].0.end > pair[1].0.start) + { + return Err(rewrite_limit_error()); + } + let mut output = formula.to_string(); + for (range, replacement) in replacements.into_iter().rev() { + if !output.is_char_boundary(range.start) || !output.is_char_boundary(range.end) { + return Err(rewrite_limit_error()); + } + output.replace_range(range, &replacement); + } + Ok(output) +} + +fn table_referenced(table: &str, reference: &str) -> UseError { + office_error( + "use.office.spreadsheet_table_referenced", + format!( + "Spreadsheet table '{table}' cannot be removed while structured reference '{reference}' still targets it." + ), + ) + .with_detail("table", table) + .with_detail("reference", reference) +} + +fn rewrite_unsupported(table: &str, reason: impl Into) -> UseError { + office_error( + "use.office.spreadsheet_table_formula_rewrite_unsupported", + format!( + "Spreadsheet table '{table}' cannot be changed without an unsafe structured-reference rewrite: {}", + reason.into() + ), + ) + .with_detail("table", table) +} + +fn rewrite_limit_error() -> UseError { + office_error( + "use.office.spreadsheet_table_formula_rewrite_unsupported", + "Spreadsheet structured-reference rewrite exceeded safe text boundaries.", + ) +} + +#[cfg(test)] +mod tests { + use super::*; + + fn rename_plan(geometry_changed: bool) -> StructuredReferenceRewritePlan { + StructuredReferenceRewritePlan::rename( + "Sales", + "Sheet1", + BTreeMap::from([("sales".into(), "Orders".into())]), + BTreeMap::from([("qty".into(), Some("Units".into()))]), + geometry_changed, + ) + } + + #[test] + fn rewrite_preserves_strings_and_external_workbooks() { + let rewritten = rename_plan(false) + .rewrite( + r#"=CONCAT("Sales[Qty]",SUM(Sales[[#Data],[Qty]]),'Sheet1'!Sales[@Qty],'[Book.xlsx]Sheet1'!Sales[Qty])"#, + LocalStructuredReferenceContext::Unknown, + ) + .unwrap(); + assert_eq!( + rewritten.formula, + r#"=CONCAT("Sales[Qty]",SUM(Orders[[#Data],[Units]]),'Sheet1'!Orders[@Units],'[Book.xlsx]Sheet1'!Sales[Qty])"# + ); + assert!(rewritten.matched); + } + + #[test] + fn rewrite_requires_provable_local_context() { + let plan = rename_plan(false); + let applied = plan + .rewrite("[@Qty]", LocalStructuredReferenceContext::Applies) + .unwrap(); + assert_eq!(applied.formula, "[@Units]"); + assert!(applied.matched); + + let unrelated = plan + .rewrite("[@Qty]", LocalStructuredReferenceContext::DoesNotApply) + .unwrap(); + assert_eq!(unrelated.formula, "[@Qty]"); + assert!(!unrelated.matched); + + assert_eq!( + plan.rewrite("[@Qty]", LocalStructuredReferenceContext::Unknown,) + .unwrap_err() + .code, + "use.office.spreadsheet_table_formula_rewrite_unsupported" + ); + } + + #[test] + fn removal_and_geometry_changes_fail_closed_for_matching_references() { + let removal = + StructuredReferenceRewritePlan::removal("Sales", "Sheet1", ["Sales".to_string()]); + assert_eq!( + removal + .rewrite( + "SUM(Sales[Qty])", + LocalStructuredReferenceContext::DoesNotApply, + ) + .unwrap_err() + .code, + "use.office.spreadsheet_table_referenced" + ); + let external = removal + .rewrite( + "'[Book.xlsx]Sheet1'!Sales[Qty]", + LocalStructuredReferenceContext::Unknown, + ) + .unwrap(); + assert_eq!(external.formula, "'[Book.xlsx]Sheet1'!Sales[Qty]"); + assert!(!external.matched); + + assert_eq!( + rename_plan(true) + .rewrite("[@Qty]", LocalStructuredReferenceContext::Applies,) + .unwrap_err() + .code, + "use.office.spreadsheet_table_formula_rewrite_unsupported" + ); + } +} diff --git a/crates/office/src/spreadsheet_formula/value.rs b/crates/office/src/spreadsheet_formula/value.rs new file mode 100644 index 00000000..de735b68 --- /dev/null +++ b/crates/office/src/spreadsheet_formula/value.rs @@ -0,0 +1,58 @@ +use serde::{Deserialize, Serialize}; + +use super::{SpreadsheetFormulaCell, SpreadsheetFormulaErrorLiteral}; + +/// Scalar or rectangular dynamic-array result produced by native calculation. +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +#[serde(tag = "type", rename_all = "kebab-case", deny_unknown_fields)] +pub enum SpreadsheetFormulaValue { + Blank, + Number { + value: String, + }, + Text { + value: String, + }, + Boolean { + value: bool, + }, + Error { + error: SpreadsheetFormulaErrorLiteral, + }, + Array { + rows: Vec>, + }, +} + +impl SpreadsheetFormulaValue { + pub fn error(error: SpreadsheetFormulaErrorLiteral) -> Self { + Self::Error { error } + } + + pub fn error_literal(&self) -> Option<&'static str> { + match self { + Self::Error { error } => Some(error.as_str()), + _ => None, + } + } +} + +/// One calculated formula anchor and its optional spill extent. +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +#[serde(rename_all = "camelCase")] +pub struct SpreadsheetFormulaCalculatedCell { + pub cell: SpreadsheetFormulaCell, + pub value: SpreadsheetFormulaValue, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub spill_range: Option, +} + +/// Deterministic read-only result from one native workbook calculation pass. +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +#[serde(rename_all = "camelCase")] +pub struct SpreadsheetFormulaCalculation { + pub formula_count: usize, + pub spill_cell_count: usize, + pub calculation_order: Vec, + pub cells: Vec, +} diff --git a/crates/office/src/spreadsheet_formula_calculation_tests.rs b/crates/office/src/spreadsheet_formula_calculation_tests.rs new file mode 100644 index 00000000..8d25a3d2 --- /dev/null +++ b/crates/office/src/spreadsheet_formula_calculation_tests.rs @@ -0,0 +1,2 @@ +include!("spreadsheet_formula_calculation_tests/calculation.rs"); +include!("spreadsheet_formula_calculation_tests/recalculation.rs"); diff --git a/crates/office/src/spreadsheet_formula_calculation_tests/calculation.rs b/crates/office/src/spreadsheet_formula_calculation_tests/calculation.rs new file mode 100644 index 00000000..c427e88a --- /dev/null +++ b/crates/office/src/spreadsheet_formula_calculation_tests/calculation.rs @@ -0,0 +1,759 @@ +use crate::{ + NativeOfficeEditor, NativeOfficeMutation, NativeOfficePartType, NativeOfficeReplayArtifact, + NativeSpreadsheetConditionalFormat, NativeSpreadsheetConditionalFormatRule, + NativeSpreadsheetDataValidation, NativeSpreadsheetDataValidationType, + NativeSpreadsheetDifferentialFormat, NativeSpreadsheetNamedRange, NativeSpreadsheetTable, + SpreadsheetCellValue, SpreadsheetFormulaErrorLiteral, SpreadsheetFormulaFunctionRegistry, + SpreadsheetFormulaValue, MAX_SPREADSHEET_FORMULA_CALCULATION_TEXT_BYTES, + MAX_SPREADSHEET_FORMULA_TEXT_BYTES, +}; + +fn number(value: &str) -> SpreadsheetCellValue { + SpreadsheetCellValue::Number { + value: value.to_string(), + } +} + +fn text(value: &str) -> SpreadsheetCellValue { + SpreadsheetCellValue::Text { + value: value.to_string(), + } +} + +fn boolean(value: bool) -> SpreadsheetCellValue { + SpreadsheetCellValue::Boolean { value } +} + +fn formula(expression: &str) -> SpreadsheetCellValue { + SpreadsheetCellValue::Formula { + expression: expression.to_string(), + } +} + +#[tokio::test] +async fn calculation_evaluates_dependency_order_operators_and_typed_functions() { + let temp = tempfile::tempdir().unwrap(); + let mut editor = NativeOfficeEditor::create(temp.path().join("calculate.xlsx")) + .await + .unwrap(); + editor.set_cell_value("/Sheet1/A1", number("2")).unwrap(); + editor + .set_cell_value("/Sheet1/A2", text("ignored")) + .unwrap(); + editor + .set_cell_value("/Sheet1/B1", formula("A1*3")) + .unwrap(); + editor + .set_cell_value("/Sheet1/C1", formula("SUM(A1:B1)+ROUND(2.55,1)")) + .unwrap(); + editor + .set_cell_value("/Sheet1/D1", formula("IF(C1=10.6,\"yes\",\"no\")")) + .unwrap(); + editor.set_cell_value("/Sheet1/E1", formula("1/0")).unwrap(); + + let calculation = editor + .snapshot() + .unwrap() + .calculate_spreadsheet_formulas() + .unwrap(); + assert_eq!(calculation.formula_count, 4); + assert_eq!(calculation.spill_cell_count, 0); + assert_eq!( + calculation + .cells + .iter() + .map(|cell| (cell.cell.path(), cell.value.clone())) + .collect::>(), + [ + ( + "/Sheet1/B1".to_string(), + SpreadsheetFormulaValue::Number { value: "6".into() } + ), + ( + "/Sheet1/C1".to_string(), + SpreadsheetFormulaValue::Number { + value: "10.6".into() + } + ), + ( + "/Sheet1/D1".to_string(), + SpreadsheetFormulaValue::Text { + value: "yes".into() + } + ), + ( + "/Sheet1/E1".to_string(), + SpreadsheetFormulaValue::error(SpreadsheetFormulaErrorLiteral::DivisionByZero) + ), + ] + ); +} + +#[tokio::test] +async fn calculation_covers_the_closed_scalar_function_registry() { + let temp = tempfile::tempdir().unwrap(); + let mut editor = NativeOfficeEditor::create(temp.path().join("functions.xlsx")) + .await + .unwrap(); + editor.set_cell_value("/Sheet1/A1", number("2")).unwrap(); + editor + .set_cell_value("/Sheet1/A2", text("ignored")) + .unwrap(); + editor.set_cell_value("/Sheet1/A3", boolean(true)).unwrap(); + for (path, expression) in [ + ("/Sheet1/B1", "SUM(A1:A3)"), + ("/Sheet1/B2", "AVERAGE(A1:A3)"), + ("/Sheet1/B3", "MIN(A1:A3)"), + ("/Sheet1/B4", "MAX(A1:A3)"), + ("/Sheet1/B5", "COUNT(\"3\",TRUE,A1:A3)"), + ("/Sheet1/B6", "COUNTA(A1:A3)"), + ("/Sheet1/B7", "ABS(-2)"), + ("/Sheet1/B8", "SQRT(9)"), + ("/Sheet1/B9", "POWER(2,3)"), + ("/Sheet1/B10", "MOD(-3,2)"), + ("/Sheet1/B11", "ROUND(-2.55,1)"), + ("/Sheet1/B12", "IFERROR(1/0,7)"), + ("/Sheet1/B13", "AND(A1:A3)"), + ("/Sheet1/B14", "OR(0,A2:A2)"), + ("/Sheet1/B15", "NOT(\"TRUE\")"), + ("/Sheet1/B16", "CONCAT(\"v=\",A1)"), + ("/Sheet1/B17", "ROW(A3)"), + ("/Sheet1/B18", "COLUMN(B1)"), + ("/Sheet1/B19", "IF(TRUE,,9)"), + ("/Sheet1/B20", "IF(FALSE,9)"), + ("/Sheet1/B21", "NA()"), + ("/Sheet1/B22", "PI()"), + ] { + editor.set_cell_value(path, formula(expression)).unwrap(); + } + + let calculation = editor + .snapshot() + .unwrap() + .calculate_spreadsheet_formulas() + .unwrap(); + let values = calculation + .cells + .iter() + .map(|cell| (cell.cell.path(), cell.value.clone())) + .collect::>(); + for (path, expected) in [ + ("/Sheet1/B1", "2"), + ("/Sheet1/B2", "2"), + ("/Sheet1/B3", "2"), + ("/Sheet1/B4", "2"), + ("/Sheet1/B5", "3"), + ("/Sheet1/B6", "3"), + ("/Sheet1/B7", "2"), + ("/Sheet1/B8", "3"), + ("/Sheet1/B9", "8"), + ("/Sheet1/B10", "1"), + ("/Sheet1/B11", "-2.6"), + ("/Sheet1/B12", "7"), + ("/Sheet1/B17", "3"), + ("/Sheet1/B18", "2"), + ] { + assert_eq!( + values[path], + SpreadsheetFormulaValue::Number { + value: expected.into() + }, + "{path}" + ); + } + assert_eq!( + values["/Sheet1/B13"], + SpreadsheetFormulaValue::Boolean { value: true } + ); + assert_eq!( + values["/Sheet1/B14"], + SpreadsheetFormulaValue::Boolean { value: false } + ); + assert_eq!( + values["/Sheet1/B15"], + SpreadsheetFormulaValue::Boolean { value: false } + ); + assert_eq!( + values["/Sheet1/B16"], + SpreadsheetFormulaValue::Text { + value: "v=2".into() + } + ); + assert_eq!(values["/Sheet1/B19"], SpreadsheetFormulaValue::Blank); + assert_eq!( + values["/Sheet1/B20"], + SpreadsheetFormulaValue::Boolean { value: false } + ); + assert_eq!( + values["/Sheet1/B21"], + SpreadsheetFormulaValue::error(SpreadsheetFormulaErrorLiteral::NotAvailable) + ); + let SpreadsheetFormulaValue::Number { value } = &values["/Sheet1/B22"] else { + panic!("PI did not return a number"); + }; + assert!((value.parse::().unwrap() - std::f64::consts::PI).abs() < f64::EPSILON); +} + +#[tokio::test] +async fn calculation_plans_dynamic_array_spills_and_detects_obstructions() { + let temp = tempfile::tempdir().unwrap(); + let mut editor = NativeOfficeEditor::create(temp.path().join("spill.xlsx")) + .await + .unwrap(); + editor + .set_cell_value("/Sheet1/A1", formula("SEQUENCE(2,3,1,1)")) + .unwrap(); + editor + .set_cell_value("/Sheet1/E1", formula("TRANSPOSE(A1#)")) + .unwrap(); + editor + .set_cell_value("/Sheet1/F2", text("blocked")) + .unwrap(); + + let calculation = editor + .snapshot() + .unwrap() + .calculate_spreadsheet_formulas() + .unwrap(); + assert_eq!(calculation.spill_cell_count, 5); + assert_eq!(calculation.cells[0].spill_range.as_deref(), Some("A1:C2")); + assert_eq!( + calculation.cells[0].value, + SpreadsheetFormulaValue::Array { + rows: vec![ + vec![ + SpreadsheetFormulaValue::Number { value: "1".into() }, + SpreadsheetFormulaValue::Number { value: "2".into() }, + SpreadsheetFormulaValue::Number { value: "3".into() }, + ], + vec![ + SpreadsheetFormulaValue::Number { value: "4".into() }, + SpreadsheetFormulaValue::Number { value: "5".into() }, + SpreadsheetFormulaValue::Number { value: "6".into() }, + ], + ] + } + ); + assert_eq!( + calculation.cells[1].value, + SpreadsheetFormulaValue::error(SpreadsheetFormulaErrorLiteral::Spill) + ); + assert_eq!(calculation.cells[1].spill_range, None); +} + +#[tokio::test] +async fn calculation_resolves_table_items_current_rows_and_local_references() { + let temp = tempfile::tempdir().unwrap(); + let mut editor = NativeOfficeEditor::create(temp.path().join("table-references.xlsx")) + .await + .unwrap(); + editor + .add_spreadsheet_table( + "/Sheet1", + NativeSpreadsheetTable::new("Sales", "A1:C5", ["Item", "Qty", "Unit Price"]) + .with_display_name("SalesView") + .with_totals_row(true), + ) + .unwrap(); + for (path, value) in [ + ("/Sheet1/B2", number("2")), + ("/Sheet1/C2", number("10")), + ("/Sheet1/B3", number("3")), + ("/Sheet1/C3", number("20")), + ("/Sheet1/C4", number("30")), + ("/Sheet1/B5", number("999")), + ("/Sheet1/C5", number("999")), + ] { + editor.set_cell_value(path, value).unwrap(); + } + editor + .set_cell_value("/Sheet1/B4", formula("B2+B3")) + .unwrap(); + editor + .set_cell_value("/Sheet1/A2", formula("Sales[@Qty]")) + .unwrap(); + editor + .set_cell_value("/Sheet1/A3", formula("SUM(Sales[[#This Row],[Qty]])")) + .unwrap(); + editor + .set_cell_value("/Sheet1/A4", formula("[@Qty]")) + .unwrap(); + editor + .set_cell_value("/Sheet1/E1", formula("SUM(SalesView[Qty])")) + .unwrap(); + editor + .set_cell_value("/Sheet1/F1", formula("SUM(Sales[[Qty]:[Unit Price]])")) + .unwrap(); + editor + .set_cell_value("/Sheet1/G1", formula("SUM(Sales[#All])")) + .unwrap(); + editor + .set_cell_value("/Sheet1/H1", formula("SUM(Sales[#Data])")) + .unwrap(); + editor + .set_cell_value("/Sheet1/I1", formula("COUNTA(Sales[[#Headers],[Qty]])")) + .unwrap(); + editor + .set_cell_value("/Sheet1/J1", formula("SUM(Sales[[#Totals],[Qty]])")) + .unwrap(); + editor + .set_cell_value( + "/Sheet1/K1", + formula("SUM(Sales[[#Headers],[#Totals],[Qty]])"), + ) + .unwrap(); + + let graph = editor + .snapshot() + .unwrap() + .formula_dependency_graph() + .unwrap(); + for path in [ + "/Sheet1/A2", + "/Sheet1/A3", + "/Sheet1/A4", + "/Sheet1/E1", + "/Sheet1/F1", + "/Sheet1/G1", + "/Sheet1/H1", + "/Sheet1/I1", + "/Sheet1/J1", + "/Sheet1/K1", + ] { + let node = graph + .nodes + .iter() + .find(|node| node.cell.path() == path) + .unwrap(); + assert!(node.unresolved_references.is_empty(), "{path}"); + } + for path in ["/Sheet1/A4", "/Sheet1/E1", "/Sheet1/F1"] { + let node = graph + .nodes + .iter() + .find(|node| node.cell.path() == path) + .unwrap(); + assert_eq!( + node.dependencies + .iter() + .map(|cell| cell.path()) + .collect::>(), + ["/Sheet1/B4"] + ); + } + for path in ["/Sheet1/G1", "/Sheet1/H1"] { + let node = graph + .nodes + .iter() + .find(|node| node.cell.path() == path) + .unwrap(); + assert_eq!( + node.dependencies + .iter() + .map(|cell| cell.path()) + .collect::>(), + ["/Sheet1/A2", "/Sheet1/A3", "/Sheet1/A4", "/Sheet1/B4"] + ); + } + + let calculation = editor.recalculate_spreadsheet_formulas().unwrap(); + let values = calculation + .cells + .iter() + .map(|cell| (cell.cell.path(), cell.value.clone())) + .collect::>(); + for (path, expected) in [ + ("/Sheet1/A2", "2"), + ("/Sheet1/A3", "3"), + ("/Sheet1/A4", "5"), + ("/Sheet1/E1", "10"), + ("/Sheet1/F1", "70"), + ("/Sheet1/G1", "2078"), + ("/Sheet1/H1", "80"), + ("/Sheet1/I1", "1"), + ("/Sheet1/J1", "999"), + ("/Sheet1/K1", "999"), + ] { + assert_eq!( + values[path], + SpreadsheetFormulaValue::Number { + value: expected.into() + }, + "{path}" + ); + } + let artifact = NativeOfficeReplayArtifact::dump(&editor.snapshot().unwrap(), "/").unwrap(); + let mut restored = + NativeOfficeEditor::create(temp.path().join("table-references-restored.xlsx")) + .await + .unwrap(); + restored.apply_replay(&artifact).unwrap(); + assert_eq!( + restored.package().content_sha256(), + editor.package().content_sha256() + ); +} + +#[tokio::test] +async fn calculation_rejects_table_local_and_missing_item_rows_atomically() { + let temp = tempfile::tempdir().unwrap(); + let mut outside = NativeOfficeEditor::create(temp.path().join("table-local-outside.xlsx")) + .await + .unwrap(); + outside + .add_spreadsheet_table( + "/Sheet1", + NativeSpreadsheetTable::new("Sales", "A1:C4", ["Item", "Qty", "Price"]), + ) + .unwrap(); + outside + .set_cell_value("/Sheet1/E2", formula("SUM([@Qty])")) + .unwrap(); + let graph = outside + .snapshot() + .unwrap() + .formula_dependency_graph() + .unwrap(); + let node = graph + .nodes + .iter() + .find(|node| node.cell.path() == "/Sheet1/E2") + .unwrap(); + assert_eq!( + node.unresolved_references[0].kind, + crate::SpreadsheetFormulaUnresolvedReferenceKind::StructuredReference + ); + let before = outside.package().content_sha256(); + let error = outside.recalculate_spreadsheet_formulas().unwrap_err(); + assert_eq!( + error.code, + "use.office.spreadsheet_formula_structured_reference_unsupported" + ); + assert_eq!(outside.package().content_sha256(), before); + + for (file_name, table, expression) in [ + ( + "table-missing-header.xlsx", + NativeSpreadsheetTable::new("Data", "A1:B3", ["Name", "Value"]).with_header_row(false), + "SUM(Data[#Headers])", + ), + ( + "table-missing-totals.xlsx", + NativeSpreadsheetTable::new("Data", "A1:B3", ["Name", "Value"]), + "SUM(Data[#Totals])", + ), + ] { + let mut editor = NativeOfficeEditor::create(temp.path().join(file_name)) + .await + .unwrap(); + editor.add_spreadsheet_table("/Sheet1", table).unwrap(); + editor + .set_cell_value("/Sheet1/D1", formula(expression)) + .unwrap(); + let before = editor.package().content_sha256(); + let error = editor.recalculate_spreadsheet_formulas().unwrap_err(); + assert_eq!( + error.code, "use.office.spreadsheet_formula_structured_reference_unsupported", + "{expression}" + ); + assert_eq!(editor.package().content_sha256(), before, "{expression}"); + } +} + +#[tokio::test] +async fn table_mutation_rewrites_structured_formula_identities_and_columns() { + let temp = tempfile::tempdir().unwrap(); + let mut editor = NativeOfficeEditor::create(temp.path().join("table-formula-rewrite.xlsx")) + .await + .unwrap(); + editor + .add_spreadsheet_table( + "/Sheet1", + NativeSpreadsheetTable::new("Sales", "A1:C4", ["Item", "Qty", "Price"]) + .with_display_name("SalesView"), + ) + .unwrap(); + editor.set_cell_value("/Sheet1/B2", number("2")).unwrap(); + editor.set_cell_value("/Sheet1/B3", number("3")).unwrap(); + editor + .set_cell_value("/Sheet1/A2", formula("[@Qty]")) + .unwrap(); + editor + .set_cell_value("/Sheet1/E1", formula("SUM(SalesView[Qty])")) + .unwrap(); + editor + .set_cell_value("/Sheet1/F1", formula("SUM(Sales[Qty])")) + .unwrap(); + editor + .set_cell_value( + "/Sheet1/G1", + formula("CONCAT(\"Sales[Qty]=\",SUM(Sales[Qty]))"), + ) + .unwrap(); + editor + .add_named_range(NativeSpreadsheetNamedRange::new( + "TableQuantity", + "SUM(Sales[Qty])", + )) + .unwrap(); + editor + .add_conditional_format( + "/Sheet1", + NativeSpreadsheetConditionalFormat::new( + "B2:B4", + NativeSpreadsheetConditionalFormatRule::Formula { + formula: "SalesView[Qty]>0".into(), + format: NativeSpreadsheetDifferentialFormat::default(), + }, + ), + ) + .unwrap(); + editor + .add_data_validation( + "/Sheet1", + NativeSpreadsheetDataValidation::new( + NativeSpreadsheetDataValidationType::Custom, + "C2:C4", + "SUM(Sales[Qty])>0", + ), + ) + .unwrap(); + editor.recalculate_spreadsheet_formulas().unwrap(); + + editor + .set_spreadsheet_table( + "/Sheet1/table[1]", + NativeSpreadsheetTable::new("Orders", "A1:C4", ["Product", "Units", "Cost"]) + .with_display_name("OrdersView"), + ) + .unwrap(); + + let snapshot = editor.snapshot().unwrap(); + for (path, expected) in [ + ("/Sheet1/A2", "[@Units]"), + ("/Sheet1/E1", "SUM(OrdersView[Units])"), + ("/Sheet1/F1", "SUM(Orders[Units])"), + ("/Sheet1/G1", "CONCAT(\"Sales[Qty]=\",SUM(Orders[Units]))"), + ] { + assert_eq!(snapshot.get(path, 0).unwrap().format["formula"], expected); + } + assert_eq!( + snapshot + .get("/namedrange[@name=TableQuantity][@scope=workbook]", 0,) + .unwrap() + .format["ref"], + "SUM(Orders[Units])" + ); + assert_eq!( + snapshot.get("/Sheet1/cf[1]", 0).unwrap().format["formula"], + "OrdersView[Units]>0" + ); + assert_eq!( + snapshot.get("/Sheet1/dataValidation[1]", 0).unwrap().format["formula1"], + "SUM(Orders[Units])>0" + ); + + let calculation = editor.recalculate_spreadsheet_formulas().unwrap(); + let values = calculation + .cells + .iter() + .map(|cell| (cell.cell.path(), cell.value.clone())) + .collect::>(); + for path in ["/Sheet1/E1", "/Sheet1/F1"] { + assert_eq!( + values[path], + SpreadsheetFormulaValue::Number { value: "5".into() }, + "{path}" + ); + } + assert_eq!( + values["/Sheet1/G1"], + SpreadsheetFormulaValue::Text { + value: "Sales[Qty]=5".into() + } + ); + + let artifact = NativeOfficeReplayArtifact::dump(&editor.snapshot().unwrap(), "/").unwrap(); + let mut restored = + NativeOfficeEditor::create(temp.path().join("table-formula-rewrite-restored.xlsx")) + .await + .unwrap(); + restored.apply_replay(&artifact).unwrap(); + assert_eq!( + restored.package().content_sha256(), + editor.package().content_sha256() + ); +} + +#[tokio::test] +async fn table_mutation_rejects_referenced_removal_and_unsafe_local_geometry_atomically() { + let temp = tempfile::tempdir().unwrap(); + let mut referenced = NativeOfficeEditor::create(temp.path().join("referenced-table.xlsx")) + .await + .unwrap(); + referenced + .add_spreadsheet_table( + "/Sheet1", + NativeSpreadsheetTable::new("Sales", "A1:C4", ["Item", "Qty", "Price"]), + ) + .unwrap(); + referenced + .set_cell_value("/Sheet1/E1", formula("SUM(Sales[Qty])")) + .unwrap(); + let before = referenced.package().content_sha256(); + let error = referenced.remove("/Sheet1/table[1]").unwrap_err(); + assert_eq!(error.code, "use.office.spreadsheet_table_referenced"); + assert_eq!(referenced.package().content_sha256(), before); + + let mut local = NativeOfficeEditor::create(temp.path().join("local-table-geometry.xlsx")) + .await + .unwrap(); + local + .add_spreadsheet_table( + "/Sheet1", + NativeSpreadsheetTable::new("Sales", "A1:C4", ["Item", "Qty", "Price"]), + ) + .unwrap(); + local + .set_cell_value("/Sheet1/A2", formula("[@Qty]")) + .unwrap(); + let before = local.package().content_sha256(); + let error = local + .set_spreadsheet_table( + "/Sheet1/table[1]", + NativeSpreadsheetTable::new("Sales", "A3:C6", ["Item", "Qty", "Price"]), + ) + .unwrap_err(); + assert_eq!( + error.code, + "use.office.spreadsheet_table_formula_rewrite_unsupported" + ); + assert_eq!(local.package().content_sha256(), before); +} + +#[tokio::test] +async fn table_mutation_rewrites_chart_and_other_table_formula_carriers() { + let temp = tempfile::tempdir().unwrap(); + let mut editor = NativeOfficeEditor::create(temp.path().join("table-formula-carriers.xlsx")) + .await + .unwrap(); + editor + .add_spreadsheet_table( + "/Sheet1", + NativeSpreadsheetTable::new("Sales", "A1:B3", ["Item", "Qty"]), + ) + .unwrap(); + editor + .add_spreadsheet_table( + "/Sheet1", + NativeSpreadsheetTable::new("Summary", "D1:D3", ["Derived"]), + ) + .unwrap(); + let chart = editor + .add_part("/Sheet1", NativeOfficePartType::Chart) + .unwrap(); + editor + .replace_xml_part( + &chart.part, + r#"Sales[Qty]1"#, + ) + .unwrap(); + + let mut package = editor.package().clone(); + let table_part = "xl/tables/table2.xml"; + let table_xml = std::str::from_utf8(package.part(table_part).unwrap()).unwrap(); + let table_xml = table_xml.replace( + r#""#, + r#"SUM(Sales[Qty])"#, + ); + assert!(table_xml.contains("calculatedColumnFormula")); + package + .set_part(table_part, table_xml.into_bytes()) + .unwrap(); + let mut editor = NativeOfficeEditor::from_package(package).unwrap(); + + editor + .set_spreadsheet_table( + "/Sheet1/table[1]", + NativeSpreadsheetTable::new("Orders", "A1:B3", ["Product", "Units"]), + ) + .unwrap(); + + let chart_xml = std::str::from_utf8( + editor + .package() + .part(chart.part.trim_start_matches('/')) + .unwrap(), + ) + .unwrap(); + assert!(chart_xml.contains("Orders[Units]")); + let table_xml = std::str::from_utf8(editor.package().part(table_part).unwrap()).unwrap(); + assert!( + table_xml.contains("SUM(Orders[Units])") + ); +} + +#[tokio::test] +async fn table_geometry_changes_clear_formula_and_chart_caches() { + let temp = tempfile::tempdir().unwrap(); + let mut editor = NativeOfficeEditor::create(temp.path().join("table-geometry-caches.xlsx")) + .await + .unwrap(); + editor + .add_spreadsheet_table( + "/Sheet1", + NativeSpreadsheetTable::new("Sales", "A1:B3", ["Item", "Qty"]), + ) + .unwrap(); + editor.set_cell_value("/Sheet1/B2", number("2")).unwrap(); + editor + .set_cell_value("/Sheet1/D1", formula("SUM(Sales[Qty])")) + .unwrap(); + editor.recalculate_spreadsheet_formulas().unwrap(); + let chart = editor + .add_part("/Sheet1", NativeOfficePartType::Chart) + .unwrap(); + editor + .replace_xml_part( + &chart.part, + r#"Sales[Qty]2"#, + ) + .unwrap(); + + editor + .set_spreadsheet_table( + "/Sheet1/table[1]", + NativeSpreadsheetTable::new("Sales", "A1:B4", ["Item", "Qty"]), + ) + .unwrap(); + + let worksheet = editor + .package() + .xml_part("xl/worksheets/sheet1.xml") + .unwrap(); + let worksheet = crate::xml_edit::index_xml(&worksheet).unwrap(); + let mut cells = Vec::new(); + worksheet.descendants_named("c", &mut cells); + let formula_cell = cells + .into_iter() + .find(|cell| cell.attributes.get("r").map(String::as_str) == Some("D1")) + .unwrap(); + assert!(formula_cell + .children + .iter() + .any(|child| child.local_name == "f")); + assert!(!formula_cell + .children + .iter() + .any(|child| child.local_name == "v")); + + let chart_xml = std::str::from_utf8( + editor + .package() + .part(chart.part.trim_start_matches('/')) + .unwrap(), + ) + .unwrap(); + assert!(chart_xml.contains("Sales[Qty]")); + assert!(!chart_xml.contains("numCache")); +} diff --git a/crates/office/src/spreadsheet_formula_calculation_tests/recalculation.rs b/crates/office/src/spreadsheet_formula_calculation_tests/recalculation.rs new file mode 100644 index 00000000..a89fa1dc --- /dev/null +++ b/crates/office/src/spreadsheet_formula_calculation_tests/recalculation.rs @@ -0,0 +1,547 @@ +#[tokio::test] +async fn calculation_treats_an_explicit_empty_string_cell_as_a_spill_obstruction() { + let temp = tempfile::tempdir().unwrap(); + let mut editor = NativeOfficeEditor::create(temp.path().join("empty-obstruction.xlsx")) + .await + .unwrap(); + editor + .set_cell_value("/Sheet1/A1", formula("SEQUENCE(1,2,1,1)")) + .unwrap(); + editor.set_cell_value("/Sheet1/B1", text("")).unwrap(); + + let calculation = editor + .snapshot() + .unwrap() + .calculate_spreadsheet_formulas() + .unwrap(); + assert_eq!( + calculation.cells[0].value, + SpreadsheetFormulaValue::error(SpreadsheetFormulaErrorLiteral::Spill) + ); + assert_eq!(calculation.cells[0].spill_range, None); +} + +#[tokio::test] +async fn calculation_rejects_cycles_and_unregistered_functions_without_mutation() { + let temp = tempfile::tempdir().unwrap(); + let mut cycle = NativeOfficeEditor::create(temp.path().join("cycle.xlsx")) + .await + .unwrap(); + cycle.set_cell_value("/Sheet1/A1", formula("B1+1")).unwrap(); + cycle.set_cell_value("/Sheet1/B1", formula("A1+1")).unwrap(); + assert_eq!( + cycle + .snapshot() + .unwrap() + .calculate_spreadsheet_formulas() + .unwrap_err() + .code, + "use.office.spreadsheet_formula_cycle" + ); + + let mut unsupported = NativeOfficeEditor::create(temp.path().join("unsupported.xlsx")) + .await + .unwrap(); + unsupported + .set_cell_value("/Sheet1/A1", formula("SHELL(\"echo unsafe\")")) + .unwrap(); + assert_eq!( + unsupported + .snapshot() + .unwrap() + .calculate_spreadsheet_formulas() + .unwrap_err() + .code, + "use.office.spreadsheet_formula_function_unsupported" + ); +} + +#[tokio::test] +async fn calculation_rejects_array_broadcasts_before_oversized_allocation() { + let temp = tempfile::tempdir().unwrap(); + let mut editor = NativeOfficeEditor::create(temp.path().join("array-limit.xlsx")) + .await + .unwrap(); + editor + .set_cell_value( + "/Sheet1/A1", + formula("SEQUENCE(100000,1)+TRANSPOSE(SEQUENCE(100000,1))"), + ) + .unwrap(); + let error = editor + .snapshot() + .unwrap() + .calculate_spreadsheet_formulas() + .unwrap_err(); + assert_eq!(error.code, "use.office.spreadsheet_formula_spill_limit"); + assert_eq!(error.details["cells"], 10_000_000_000_u64); +} + +#[tokio::test] +async fn calculation_bounds_cumulative_function_argument_arrays() { + let temp = tempfile::tempdir().unwrap(); + let mut editor = NativeOfficeEditor::create(temp.path().join("function-argument-limit.xlsx")) + .await + .unwrap(); + editor + .set_cell_value( + "/Sheet1/A1", + formula("SUM(SEQUENCE(60000,1),SEQUENCE(60000,1))"), + ) + .unwrap(); + let error = editor + .snapshot() + .unwrap() + .calculate_spreadsheet_formulas() + .unwrap_err(); + assert_eq!( + error.code, + "use.office.spreadsheet_formula_function_array_limit" + ); + assert_eq!(error.details["cells"], 120_000); +} + +#[tokio::test] +async fn calculation_bounds_cumulative_spill_cells() { + let temp = tempfile::tempdir().unwrap(); + let mut editor = NativeOfficeEditor::create(temp.path().join("spill-total-limit.xlsx")) + .await + .unwrap(); + editor + .set_cell_value("/Sheet1/A1", formula("SEQUENCE(60000,1)")) + .unwrap(); + editor + .set_cell_value("/Sheet1/C1", formula("SEQUENCE(60000,1)")) + .unwrap(); + let error = editor + .snapshot() + .unwrap() + .calculate_spreadsheet_formulas() + .unwrap_err(); + assert_eq!(error.code, "use.office.spreadsheet_formula_spill_limit"); + assert_eq!(error.details["cells"], 119_998); +} + +#[tokio::test] +async fn calculation_bounds_concatenated_text_results() { + let temp = tempfile::tempdir().unwrap(); + let mut editor = NativeOfficeEditor::create(temp.path().join("text-limit.xlsx")) + .await + .unwrap(); + let chunk = "x".repeat(600_000); + editor.set_cell_value("/Sheet1/A1", text(&chunk)).unwrap(); + editor.set_cell_value("/Sheet1/A2", text(&chunk)).unwrap(); + editor + .set_cell_value("/Sheet1/A3", formula("A1&A2")) + .unwrap(); + let error = editor + .snapshot() + .unwrap() + .calculate_spreadsheet_formulas() + .unwrap_err(); + assert_eq!(error.code, "use.office.spreadsheet_formula_text_limit"); + assert_eq!(error.details["bytes"], 1_200_000); +} + +#[tokio::test] +async fn calculation_bounds_passthrough_text_results() { + let temp = tempfile::tempdir().unwrap(); + let mut editor = NativeOfficeEditor::create(temp.path().join("passthrough-text-limit.xlsx")) + .await + .unwrap(); + let oversized = "x".repeat(MAX_SPREADSHEET_FORMULA_TEXT_BYTES + 1); + editor + .set_cell_value("/Sheet1/A1", text(&oversized)) + .unwrap(); + editor.set_cell_value("/Sheet1/A2", formula("A1")).unwrap(); + let error = editor + .snapshot() + .unwrap() + .calculate_spreadsheet_formulas() + .unwrap_err(); + assert_eq!(error.code, "use.office.spreadsheet_formula_text_limit"); + assert_eq!( + error.details["bytes"], + MAX_SPREADSHEET_FORMULA_TEXT_BYTES + 1 + ); +} + +#[tokio::test] +async fn calculation_bounds_cumulative_text_results() { + let temp = tempfile::tempdir().unwrap(); + let mut editor = NativeOfficeEditor::create(temp.path().join("text-total-limit.xlsx")) + .await + .unwrap(); + let text_result = "x".repeat(MAX_SPREADSHEET_FORMULA_TEXT_BYTES); + editor + .set_cell_value("/Sheet1/A1", text(&text_result)) + .unwrap(); + for row in 1..=9 { + editor + .set_cell_value(format!("/Sheet1/B{row}"), formula("A1")) + .unwrap(); + } + let error = editor + .snapshot() + .unwrap() + .calculate_spreadsheet_formulas() + .unwrap_err(); + assert_eq!(error.code, "use.office.spreadsheet_formula_text_limit"); + assert_eq!( + error.details["bytes"], + MAX_SPREADSHEET_FORMULA_CALCULATION_TEXT_BYTES + MAX_SPREADSHEET_FORMULA_TEXT_BYTES + ); +} + +#[tokio::test] +async fn calculation_bounds_broadcast_text_before_oversized_array_allocation() { + let temp = tempfile::tempdir().unwrap(); + let mut editor = NativeOfficeEditor::create(temp.path().join("text-array-limit.xlsx")) + .await + .unwrap(); + let text_result = "x".repeat(MAX_SPREADSHEET_FORMULA_TEXT_BYTES); + editor + .set_cell_value("/Sheet1/A1", text(&text_result)) + .unwrap(); + editor + .set_cell_value("/Sheet1/B1", formula("IF(SEQUENCE(9,1)>0,A1,\"\")")) + .unwrap(); + let error = editor + .snapshot() + .unwrap() + .calculate_spreadsheet_formulas() + .unwrap_err(); + assert_eq!(error.code, "use.office.spreadsheet_formula_text_limit"); + assert_eq!( + error.details["bytes"], + MAX_SPREADSHEET_FORMULA_CALCULATION_TEXT_BYTES + MAX_SPREADSHEET_FORMULA_TEXT_BYTES + ); +} + +#[tokio::test] +async fn recalculation_atomically_writes_cached_values_spills_and_calculation_metadata() { + let temp = tempfile::tempdir().unwrap(); + let mut editor = NativeOfficeEditor::create(temp.path().join("recalculate.xlsx")) + .await + .unwrap(); + editor.set_cell_value("/Sheet1/A1", number("2")).unwrap(); + editor + .set_cell_value("/Sheet1/B1", formula("A1*3")) + .unwrap(); + editor + .set_cell_value("/Sheet1/C1", formula("IF(B1=6,\"yes\",\"no\")")) + .unwrap(); + editor.set_cell_value("/Sheet1/D1", formula("1/0")).unwrap(); + editor + .set_cell_value("/Sheet1/E1", formula("SEQUENCE(2,2,1,1)")) + .unwrap(); + + let calculation = editor.recalculate_spreadsheet_formulas().unwrap(); + assert_eq!(calculation.formula_count, 4); + assert_eq!(calculation.spill_cell_count, 3); + + let document = editor.snapshot().unwrap(); + let b1 = document.get("/Sheet1/B1", 0).unwrap(); + assert_eq!(b1.text, "6"); + assert_eq!(b1.format.get("formula").map(String::as_str), Some("A1*3")); + let c1 = document.get("/Sheet1/C1", 0).unwrap(); + assert_eq!(c1.text, "yes"); + assert_eq!( + c1.format.get("valueType").map(String::as_str), + Some("String") + ); + let d1 = document.get("/Sheet1/D1", 0).unwrap(); + assert_eq!(d1.text, "#DIV/0!"); + assert_eq!( + d1.format.get("valueType").map(String::as_str), + Some("Error") + ); + let e1 = document.get("/Sheet1/E1", 0).unwrap(); + assert_eq!(e1.text, "1"); + assert_eq!( + e1.format.get("formulaType").map(String::as_str), + Some("array") + ); + assert_eq!( + e1.format.get("formulaRef").map(String::as_str), + Some("E1:F2") + ); + for (path, expected) in [ + ("/Sheet1/F1", "2"), + ("/Sheet1/E2", "3"), + ("/Sheet1/F2", "4"), + ] { + let cell = document.get(path, 0).unwrap(); + assert_eq!(cell.text, expected); + assert!(!cell.format.contains_key("formula")); + } + let workbook = + String::from_utf8(editor.package().part("xl/workbook.xml").unwrap().to_vec()).unwrap(); + assert!(workbook.contains("calcCompleted=\"1\""), "{workbook}"); + assert!(workbook.contains("forceFullCalc=\"0\""), "{workbook}"); + assert!(workbook.contains("fullCalcOnLoad=\"0\""), "{workbook}"); + assert!(!editor.package().contains_part("xl/calcChain.xml")); +} + +#[tokio::test] +async fn recalculation_preserves_explicit_normal_formula_storage() { + let temp = tempfile::tempdir().unwrap(); + let mut editor = NativeOfficeEditor::create(temp.path().join("normal-formula.xlsx")) + .await + .unwrap(); + editor.set_cell_value("/Sheet1/A1", number("2")).unwrap(); + editor + .set_cell_value("/Sheet1/B1", formula("A1*3")) + .unwrap(); + + let mut package = editor.package().clone(); + let worksheet = String::from_utf8(package.part("xl/worksheets/sheet1.xml").unwrap().to_vec()) + .unwrap() + .replace("A1*3", "A1*3"); + assert!(worksheet.contains("A1*3")); + package + .set_part("xl/worksheets/sheet1.xml", worksheet.into_bytes()) + .unwrap(); + let mut editor = NativeOfficeEditor::from_package(package).unwrap(); + + editor.recalculate_spreadsheet_formulas().unwrap(); + let worksheet = String::from_utf8( + editor + .package() + .part("xl/worksheets/sheet1.xml") + .unwrap() + .to_vec(), + ) + .unwrap(); + assert!(worksheet.contains("A1*36")); +} + +#[tokio::test] +async fn recalculation_rejects_malformed_formula_storage_without_mutation() { + let temp = tempfile::tempdir().unwrap(); + let mut source = NativeOfficeEditor::create(temp.path().join("malformed-storage.xlsx")) + .await + .unwrap(); + source.set_cell_value("/Sheet1/A1", formula("1+1")).unwrap(); + + for replacement in [ + "1+1", + "1+1", + ] { + let mut package = source.package().clone(); + let worksheet = + String::from_utf8(package.part("xl/worksheets/sheet1.xml").unwrap().to_vec()) + .unwrap() + .replace("1+1", replacement); + assert!(worksheet.contains(replacement)); + package + .set_part("xl/worksheets/sheet1.xml", worksheet.into_bytes()) + .unwrap(); + let mut editor = NativeOfficeEditor::from_package(package).unwrap(); + let before = editor.package().content_sha256(); + + let error = editor.recalculate_spreadsheet_formulas().unwrap_err(); + assert_eq!(error.code, "use.office.spreadsheet_formula_storage_invalid"); + assert_eq!(editor.package().content_sha256(), before); + } +} + +#[tokio::test] +async fn recalculation_clears_cells_from_a_previous_larger_spill() { + let temp = tempfile::tempdir().unwrap(); + let mut editor = NativeOfficeEditor::create(temp.path().join("shrink-spill.xlsx")) + .await + .unwrap(); + editor.set_cell_value("/Sheet1/D1", number("2")).unwrap(); + editor + .set_cell_value("/Sheet1/A1", formula("SEQUENCE(D1,2,1,1)")) + .unwrap(); + editor.recalculate_spreadsheet_formulas().unwrap(); + assert!(editor.snapshot().unwrap().get("/Sheet1/B2", 0).is_ok()); + + editor.set_cell_value("/Sheet1/D1", number("1")).unwrap(); + let calculation = editor.recalculate_spreadsheet_formulas().unwrap(); + assert_eq!(calculation.cells[0].spill_range.as_deref(), Some("A1:B1")); + let document = editor.snapshot().unwrap(); + assert!(document.get("/Sheet1/A2", 0).is_err()); + assert!(document.get("/Sheet1/B2", 0).is_err()); + let worksheet = String::from_utf8( + editor + .package() + .part("xl/worksheets/sheet1.xml") + .unwrap() + .to_vec(), + ) + .unwrap(); + assert!(!worksheet.contains("r=\"A2\""), "{worksheet}"); + assert!(!worksheet.contains("r=\"B2\""), "{worksheet}"); +} + +#[tokio::test] +async fn spilled_cells_are_read_only_and_replacing_the_anchor_clears_the_spill() { + let temp = tempfile::tempdir().unwrap(); + let mut editor = NativeOfficeEditor::create(temp.path().join("spill-edit.xlsx")) + .await + .unwrap(); + editor + .set_cell_value("/Sheet1/A1", formula("SEQUENCE(2,2,1,1)")) + .unwrap(); + editor.recalculate_spreadsheet_formulas().unwrap(); + + let before = editor.package().content_sha256(); + let error = editor + .set_cell_value("/Sheet1/B2", text("not allowed")) + .unwrap_err(); + assert_eq!( + error.code, + "use.office.spreadsheet_formula_spill_cell_read_only" + ); + assert_eq!(editor.package().content_sha256(), before); + + editor.set_cell_value("/Sheet1/A1", number("9")).unwrap(); + let document = editor.snapshot().unwrap(); + assert_eq!(document.get("/Sheet1/A1", 0).unwrap().text, "9"); + for path in ["/Sheet1/B1", "/Sheet1/A2", "/Sheet1/B2"] { + assert!(document.get(path, 0).is_err(), "{path}"); + } +} + +#[tokio::test] +async fn removing_a_formula_anchor_removes_its_spilled_cells() { + let temp = tempfile::tempdir().unwrap(); + let mut editor = NativeOfficeEditor::create(temp.path().join("spill-remove.xlsx")) + .await + .unwrap(); + editor + .set_cell_value("/Sheet1/C3", formula("SEQUENCE(2,2,1,1)")) + .unwrap(); + editor.recalculate_spreadsheet_formulas().unwrap(); + editor.remove("/Sheet1/C3").unwrap(); + + let document = editor.snapshot().unwrap(); + for path in ["/Sheet1/C3", "/Sheet1/D3", "/Sheet1/C4", "/Sheet1/D4"] { + assert!(document.get(path, 0).is_err(), "{path}"); + } +} + +#[tokio::test] +async fn recalculation_mutation_rolls_back_the_entire_batch_on_failure() { + let temp = tempfile::tempdir().unwrap(); + let mut editor = NativeOfficeEditor::create(temp.path().join("rollback.xlsx")) + .await + .unwrap(); + editor + .set_cell_value("/Sheet1/A1", formula("SHELL(\"unsafe\")")) + .unwrap(); + let before = editor.package().content_sha256(); + let error = editor + .apply_batch(&[ + NativeOfficeMutation::SetCellValue { + path: "/Sheet1/B1".into(), + value: text("must roll back"), + }, + NativeOfficeMutation::RecalculateSpreadsheetFormulas, + ]) + .unwrap_err(); + assert_eq!( + error.code, + "use.office.spreadsheet_formula_function_unsupported" + ); + assert_eq!(editor.package().content_sha256(), before); + assert!(editor.snapshot().unwrap().get("/Sheet1/B1", 0).is_err()); +} + +#[tokio::test] +async fn recalculated_formulas_and_spills_are_exactly_replayable() { + let temp = tempfile::tempdir().unwrap(); + let mut source = NativeOfficeEditor::create(temp.path().join("source.xlsx")) + .await + .unwrap(); + source.set_cell_value("/Sheet1/A1", number("4")).unwrap(); + source + .set_cell_value("/Sheet1/B1", formula("A1/2")) + .unwrap(); + source + .set_cell_value("/Sheet1/C1", formula("SEQUENCE(2,2,1,1)")) + .unwrap(); + source.recalculate_spreadsheet_formulas().unwrap(); + + let artifact = NativeOfficeReplayArtifact::dump(&source.snapshot().unwrap(), "/").unwrap(); + assert!(matches!( + artifact.mutations.last(), + Some(NativeOfficeMutation::RecalculateSpreadsheetFormulas) + )); + let mut restored = NativeOfficeEditor::create(temp.path().join("restored.xlsx")) + .await + .unwrap(); + restored.apply_replay(&artifact).unwrap(); + assert_eq!( + restored.package().content_sha256(), + source.package().content_sha256() + ); +} + +#[tokio::test] +async fn exact_replay_rejects_an_uncached_array_formula() { + let temp = tempfile::tempdir().unwrap(); + let mut editor = NativeOfficeEditor::create(temp.path().join("uncached-array.xlsx")) + .await + .unwrap(); + editor + .set_cell_value("/Sheet1/A1", formula("SEQUENCE(2,1)")) + .unwrap(); + + let mut package = editor.package().clone(); + let worksheet = String::from_utf8(package.part("xl/worksheets/sheet1.xml").unwrap().to_vec()) + .unwrap() + .replace( + "SEQUENCE(2,1)", + "SEQUENCE(2,1)", + ); + assert!(worksheet.contains("SEQUENCE(2,1)")); + package + .set_part("xl/worksheets/sheet1.xml", worksheet.into_bytes()) + .unwrap(); + let editor = NativeOfficeEditor::from_package(package).unwrap(); + + let error = NativeOfficeReplayArtifact::dump(&editor.snapshot().unwrap(), "/").unwrap_err(); + assert_eq!(error.code, "use.office.dump_unsupported"); + assert!(error.message.contains("no cached native result")); +} + +#[tokio::test] +async fn exact_replay_rejects_explicit_normal_formula_storage() { + let temp = tempfile::tempdir().unwrap(); + let mut editor = NativeOfficeEditor::create(temp.path().join("explicit-normal.xlsx")) + .await + .unwrap(); + editor.set_cell_value("/Sheet1/A1", formula("1+1")).unwrap(); + + let mut package = editor.package().clone(); + let worksheet = String::from_utf8(package.part("xl/worksheets/sheet1.xml").unwrap().to_vec()) + .unwrap() + .replace("1+1", "1+1"); + assert!(worksheet.contains("1+1")); + package + .set_part("xl/worksheets/sheet1.xml", worksheet.into_bytes()) + .unwrap(); + let editor = NativeOfficeEditor::from_package(package).unwrap(); + + let error = NativeOfficeReplayArtifact::dump(&editor.snapshot().unwrap(), "/").unwrap_err(); + assert_eq!(error.code, "use.office.dump_unsupported"); + assert!(error.message.contains("not canonical replay input")); +} + +#[test] +fn typed_function_registry_is_closed_and_send_sync() { + fn assert_send_sync() {} + assert_send_sync::(); + let registry = SpreadsheetFormulaFunctionRegistry::default(); + assert!(registry.contains("SUM")); + assert!(registry.contains("_xlfn.SEQUENCE")); + assert!(!registry.contains("SHELL")); + assert_eq!( + serde_json::to_value(NativeOfficeMutation::RecalculateSpreadsheetFormulas).unwrap(), + serde_json::json!({"operation": "recalculate-spreadsheet-formulas"}) + ); +} diff --git a/crates/office/src/spreadsheet_formula_graph_tests.rs b/crates/office/src/spreadsheet_formula_graph_tests.rs new file mode 100644 index 00000000..9f248fbb --- /dev/null +++ b/crates/office/src/spreadsheet_formula_graph_tests.rs @@ -0,0 +1,226 @@ +use crate::{ + NativeOfficeEditor, NativeOfficeMutation, NativeSpreadsheetNamedRange, SpreadsheetCellValue, + SpreadsheetFormulaDependencyGraph, SpreadsheetFormulaUnresolvedReferenceKind, + MAX_SPREADSHEET_FORMULA_DEPTH, MAX_SPREADSHEET_FORMULA_REFERENCE_VISITS, +}; + +fn number(value: &str) -> SpreadsheetCellValue { + SpreadsheetCellValue::Number { + value: value.to_string(), + } +} + +fn formula(expression: &str) -> SpreadsheetCellValue { + SpreadsheetCellValue::Formula { + expression: expression.to_string(), + } +} + +#[tokio::test] +async fn dependency_graph_orders_cross_sheet_ranges_and_named_references() { + let temp = tempfile::tempdir().unwrap(); + let mut editor = NativeOfficeEditor::create(temp.path().join("graph.xlsx")) + .await + .unwrap(); + editor.add_worksheet("Data").unwrap(); + editor.set_cell_value("/Sheet1/A1", number("2")).unwrap(); + editor + .set_cell_value("/Sheet1/B1", formula("A1*2")) + .unwrap(); + editor + .set_cell_value("/Data/A1", formula("Sheet1!B1+1")) + .unwrap(); + editor + .add_named_range(NativeSpreadsheetNamedRange::new( + "Inputs", + "'Sheet1'!$B$1:$B$3", + )) + .unwrap(); + editor + .set_cell_value("/Sheet1/B3", formula("B1+1")) + .unwrap(); + editor + .set_cell_value("/Sheet1/C1", formula("SUM(Inputs)+Data!A1")) + .unwrap(); + + let graph = editor + .snapshot() + .unwrap() + .formula_dependency_graph() + .unwrap(); + assert_eq!(graph.nodes.len(), 4); + assert!(graph.cycles.is_empty()); + assert_eq!( + graph + .calculation_order + .iter() + .map(|cell| cell.path()) + .collect::>(), + ["/Sheet1/B1", "/Sheet1/B3", "/Data/A1", "/Sheet1/C1"] + ); + let target = graph + .nodes + .iter() + .find(|node| node.cell.path() == "/Sheet1/C1") + .unwrap(); + assert_eq!( + target + .dependencies + .iter() + .map(|cell| cell.path()) + .collect::>(), + ["/Sheet1/B1", "/Sheet1/B3", "/Data/A1"] + ); + assert!(target.unresolved_references.is_empty()); +} + +#[tokio::test] +async fn dependency_graph_reports_stable_cycles_and_unresolved_references() { + let temp = tempfile::tempdir().unwrap(); + let mut editor = NativeOfficeEditor::create(temp.path().join("cycles.xlsx")) + .await + .unwrap(); + editor + .set_cell_value("/Sheet1/A1", formula("B1+Missing!A1")) + .unwrap(); + editor + .set_cell_value("/Sheet1/B1", formula("A1+UnknownName")) + .unwrap(); + editor + .set_cell_value("/Sheet1/C1", formula("[Book.xlsx]Data!A1")) + .unwrap(); + + let graph = editor + .snapshot() + .unwrap() + .formula_dependency_graph() + .unwrap(); + assert_eq!( + graph + .cycles + .iter() + .map(|cycle| cycle.iter().map(|cell| cell.path()).collect::>()) + .collect::>(), + [vec!["/Sheet1/A1", "/Sheet1/B1"]] + ); + assert_eq!( + graph + .nodes + .iter() + .find(|node| node.cell.path() == "/Sheet1/A1") + .unwrap() + .unresolved_references[0] + .kind, + SpreadsheetFormulaUnresolvedReferenceKind::MissingWorksheet + ); + assert_eq!( + graph + .nodes + .iter() + .find(|node| node.cell.path() == "/Sheet1/B1") + .unwrap() + .unresolved_references[0] + .kind, + SpreadsheetFormulaUnresolvedReferenceKind::UndefinedName + ); + assert_eq!( + graph + .nodes + .iter() + .find(|node| node.cell.path() == "/Sheet1/C1") + .unwrap() + .unresolved_references[0] + .kind, + SpreadsheetFormulaUnresolvedReferenceKind::ExternalWorkbook + ); +} + +#[tokio::test] +async fn dependency_graph_and_calculation_bound_nested_named_references() { + let temp = tempfile::tempdir().unwrap(); + let mut editor = NativeOfficeEditor::create(temp.path().join("named-depth.xlsx")) + .await + .unwrap(); + editor.set_cell_value("/Sheet1/B1", number("1")).unwrap(); + let mutations = (0..=MAX_SPREADSHEET_FORMULA_DEPTH) + .map(|index| { + let reference = if index == MAX_SPREADSHEET_FORMULA_DEPTH { + "'Sheet1'!$B$1".to_string() + } else { + format!("Chain_{}", index + 1) + }; + NativeOfficeMutation::AddNamedRange { + named_range: NativeSpreadsheetNamedRange::new(format!("Chain_{index}"), reference), + } + }) + .collect::>(); + editor.apply_batch(&mutations).unwrap(); + editor + .set_cell_value("/Sheet1/A1", formula("Chain_0")) + .unwrap(); + + let graph = editor + .snapshot() + .unwrap() + .formula_dependency_graph() + .unwrap(); + assert_eq!( + graph.nodes[0].unresolved_references[0].kind, + SpreadsheetFormulaUnresolvedReferenceKind::NamedRangeDepth + ); + let error = editor + .snapshot() + .unwrap() + .calculate_spreadsheet_formulas() + .unwrap_err(); + assert_eq!( + error.code, + "use.office.spreadsheet_formula_named_reference_depth" + ); + assert_eq!(error.details["namedRange"], "Chain_128"); +} + +#[tokio::test] +async fn dependency_graph_bounds_overlapping_reference_candidate_visits() { + let temp = tempfile::tempdir().unwrap(); + let mut editor = NativeOfficeEditor::create(temp.path().join("reference-visits.xlsx")) + .await + .unwrap(); + const FORMULAS: usize = 4_000; + editor + .set_cell_value( + format!("/Sheet1/A1:A{FORMULAS}"), + SpreadsheetCellValue::Formula { + expression: "1".into(), + }, + ) + .unwrap(); + let area_count = MAX_SPREADSHEET_FORMULA_REFERENCE_VISITS / FORMULAS + 1; + let areas = (0..area_count) + .map(|offset| format!("A1:A{}", FORMULAS + offset)) + .collect::>() + .join(","); + editor + .set_cell_value("/Sheet1/B1", formula(&format!("SUM({areas})"))) + .unwrap(); + + let error = editor + .snapshot() + .unwrap() + .formula_dependency_graph() + .unwrap_err(); + assert_eq!( + error.code, + "use.office.spreadsheet_formula_reference_visit_limit" + ); + assert_eq!( + error.details["visits"], + MAX_SPREADSHEET_FORMULA_REFERENCE_VISITS + 1 + ); +} + +#[test] +fn dependency_graph_contract_is_send_sync() { + fn assert_send_sync() {} + assert_send_sync::(); +} diff --git a/crates/office/src/spreadsheet_import_tests.rs b/crates/office/src/spreadsheet_import_tests.rs new file mode 100644 index 00000000..b46418bc --- /dev/null +++ b/crates/office/src/spreadsheet_import_tests.rs @@ -0,0 +1,356 @@ +use crate::{ + NativeOfficeEditor, NativeOfficeMutation, NativeOfficeReplayArtifact, + NativeSpreadsheetDelimitedFormat, NativeSpreadsheetDelimitedImport, + NativeSpreadsheetFrozenPane, OfficeNodeType, SpreadsheetCellValue, +}; + +fn text(value: &str) -> SpreadsheetCellValue { + SpreadsheetCellValue::Text { + value: value.to_string(), + } +} + +fn cell(editor: &NativeOfficeEditor, path: &str) -> crate::DocumentNode { + editor.snapshot().unwrap().get(path, 0).unwrap() +} + +#[test] +fn delimited_import_and_frozen_pane_have_closed_typed_batch_contracts() { + fn assert_send_sync() {} + assert_send_sync::(); + assert_send_sync::(); + + let mutation = NativeOfficeMutation::ImportSpreadsheetDelimited { + sheet: "/Sheet1".into(), + import: NativeSpreadsheetDelimitedImport::new( + "Name\tValue\nAlpha\t42", + NativeSpreadsheetDelimitedFormat::Tsv, + ) + .with_header(true) + .with_start_cell("B2"), + }; + assert_eq!( + serde_json::to_value(mutation).unwrap(), + serde_json::json!({ + "operation": "import-spreadsheet-delimited", + "sheet": "/Sheet1", + "import": { + "content": "Name\tValue\nAlpha\t42", + "format": "tsv", + "header": true, + "startCell": "B2" + } + }) + ); + + let pane = NativeOfficeMutation::SetSpreadsheetFrozenPane { + sheet: "/Sheet1".into(), + pane: NativeSpreadsheetFrozenPane::new(2, 0, "B3"), + }; + assert_eq!( + serde_json::to_value(pane).unwrap(), + serde_json::json!({ + "operation": "set-spreadsheet-frozen-pane", + "sheet": "/Sheet1", + "pane": { + "frozenRows": 2, + "frozenColumns": 0, + "topLeftCell": "B3" + } + }) + ); +} + +#[tokio::test] +async fn import_handles_bom_quotes_blank_rows_types_header_filter_and_freeze() { + let temp = tempfile::tempdir().unwrap(); + let path = temp.path().join("import.xlsx"); + let mut editor = NativeOfficeEditor::create(&path).await.unwrap(); + let content = "\u{feff}Name,Value,Note\r\n\"Alpha, Inc\",001,\"line 1\nline 2\"\r\n\r\nBeta,TRUE,2026-07-17"; + + let receipt = editor + .import_spreadsheet_delimited( + "/Sheet1", + NativeSpreadsheetDelimitedImport::new(content, NativeSpreadsheetDelimitedFormat::Csv) + .with_header(true) + .with_start_cell("B2"), + ) + .unwrap(); + + assert_eq!(receipt.path, "/Sheet1/B2:D5"); + assert_eq!(receipt.range.as_deref(), Some("B2:D5")); + assert_eq!(receipt.row_count, 4); + assert_eq!(receipt.column_count, 3); + assert_eq!(receipt.filter_path.as_deref(), Some("/Sheet1/autofilter")); + assert_eq!(receipt.freeze_path.as_deref(), Some("/Sheet1/freeze")); + assert!(receipt.changed); + + assert_eq!(cell(&editor, "/Sheet1/B3").text, "Alpha, Inc"); + assert_eq!(cell(&editor, "/Sheet1/C3").text, "001"); + assert_eq!(cell(&editor, "/Sheet1/C3").format["valueType"], "Number"); + assert_eq!(cell(&editor, "/Sheet1/D3").text, "line 1\nline 2"); + assert!(editor.snapshot().unwrap().get("/Sheet1/B4", 0).is_err()); + assert_eq!(cell(&editor, "/Sheet1/C5").text, "true"); + assert_eq!(cell(&editor, "/Sheet1/C5").format["valueType"], "Boolean"); + assert_eq!( + cell(&editor, "/Sheet1/D5").format["numberFormat"], + "yyyy-mm-dd" + ); + + let snapshot = editor.snapshot().unwrap(); + let filter = snapshot.get("/Sheet1/autofilter", 0).unwrap(); + assert_eq!(filter.format["ref"], "B2:D5"); + let freeze = snapshot.get("/Sheet1/freeze", 0).unwrap(); + assert_eq!(freeze.node_type, OfficeNodeType::FrozenPane); + assert_eq!(freeze.format["frozenRows"], "2"); + assert_eq!(freeze.format["frozenColumns"], "0"); + assert_eq!(freeze.format["topLeftCell"], "B3"); + assert_eq!(freeze.format["nativeMutable"], "true"); +} + +#[tokio::test] +async fn import_upserts_existing_cells_preserves_ragged_columns_and_clears_explicit_empties() { + let temp = tempfile::tempdir().unwrap(); + let path = temp.path().join("upsert.xlsx"); + let mut editor = NativeOfficeEditor::create(&path).await.unwrap(); + for (path, value) in [ + ("/Sheet1/A2", "left"), + ("/Sheet1/B2", "old-b2"), + ("/Sheet1/C2", "old-c2"), + ("/Sheet1/D2", "right-2"), + ("/Sheet1/B3", "old-b3"), + ("/Sheet1/C3", "old-c3"), + ("/Sheet1/D3", "right-3"), + ] { + editor.set_cell_value(path, text(value)).unwrap(); + } + + let result = editor + .import_spreadsheet_delimited( + "/Sheet1", + NativeSpreadsheetDelimitedImport::new( + "new-b2,\n,new-c3", + NativeSpreadsheetDelimitedFormat::Csv, + ) + .with_start_cell("B2"), + ) + .unwrap(); + assert_eq!(result.range.as_deref(), Some("B2:C3")); + assert_eq!(cell(&editor, "/Sheet1/A2").text, "left"); + assert_eq!(cell(&editor, "/Sheet1/B2").text, "new-b2"); + assert_eq!(cell(&editor, "/Sheet1/C2").text, ""); + assert_eq!(cell(&editor, "/Sheet1/B3").text, ""); + assert_eq!(cell(&editor, "/Sheet1/C3").text, "new-c3"); + assert_eq!(cell(&editor, "/Sheet1/D2").text, "right-2"); + assert_eq!(cell(&editor, "/Sheet1/D3").text, "right-3"); + + let quoted_empty = editor + .import_spreadsheet_delimited( + "/Sheet1", + NativeSpreadsheetDelimitedImport::new("\"\"", NativeSpreadsheetDelimitedFormat::Csv) + .with_start_cell("A2"), + ) + .unwrap(); + assert_eq!(quoted_empty.row_count, 1); + assert_eq!(quoted_empty.column_count, 1); + assert_eq!(cell(&editor, "/Sheet1/A2").text, ""); +} + +#[tokio::test] +async fn import_validates_before_commit_and_failed_batches_roll_back() { + let temp = tempfile::tempdir().unwrap(); + let path = temp.path().join("rollback.xlsx"); + let mut editor = NativeOfficeEditor::create(&path).await.unwrap(); + let before = editor.package().content_sha256(); + let error = editor + .apply_batch(&[ + NativeOfficeMutation::SetCellValue { + path: "/Sheet1/A1".into(), + value: text("must roll back"), + }, + NativeOfficeMutation::ImportSpreadsheetDelimited { + sheet: "/Sheet1".into(), + import: NativeSpreadsheetDelimitedImport::new( + "one,two", + NativeSpreadsheetDelimitedFormat::Csv, + ) + .with_start_cell("XFD1"), + }, + ]) + .unwrap_err(); + assert_eq!(error.code, "use.office.spreadsheet_import_column_limit"); + assert_eq!(editor.package().content_sha256(), before); + + let long = "x".repeat(32_768); + let error = editor + .import_spreadsheet_delimited( + "/Sheet1", + NativeSpreadsheetDelimitedImport::new(long, NativeSpreadsheetDelimitedFormat::Csv), + ) + .unwrap_err(); + assert_eq!(error.code, "use.office.spreadsheet_import_cell_limit"); + assert_eq!(editor.package().content_sha256(), before); + + for malformed in ["\"unclosed", "\"closed\"suffix", "unquoted\"quote"] { + let error = editor + .import_spreadsheet_delimited( + "/Sheet1", + NativeSpreadsheetDelimitedImport::new( + malformed, + NativeSpreadsheetDelimitedFormat::Csv, + ), + ) + .unwrap_err(); + assert_eq!( + error.code, + "use.office.spreadsheet_import_delimited_invalid" + ); + assert_eq!(editor.package().content_sha256(), before); + } +} + +#[tokio::test] +async fn frozen_pane_has_a_typed_semantic_lifecycle_and_atomic_validation() { + let temp = tempfile::tempdir().unwrap(); + let path = temp.path().join("freeze.xlsx"); + let mut editor = NativeOfficeEditor::create(&path).await.unwrap(); + + let path = editor + .set_spreadsheet_frozen_pane("/Sheet1", NativeSpreadsheetFrozenPane::new(1, 2, "C2")) + .unwrap(); + assert_eq!(path, "/Sheet1/freeze"); + let snapshot = editor.snapshot().unwrap(); + let pane = snapshot.get("/Sheet1/freeze", 0).unwrap(); + assert_eq!(pane.format["frozenRows"], "1"); + assert_eq!(pane.format["frozenColumns"], "2"); + assert_eq!(pane.format["topLeftCell"], "C2"); + assert_eq!(pane.format["activePane"], "bottomRight"); + assert_eq!(snapshot.query("frozen-pane").unwrap().len(), 1); + + let before = editor.package().content_sha256(); + let error = editor + .apply_batch(&[ + NativeOfficeMutation::SetCellValue { + path: "/Sheet1/A1".into(), + value: text("must roll back"), + }, + NativeOfficeMutation::SetSpreadsheetFrozenPane { + sheet: "/Sheet1".into(), + pane: NativeSpreadsheetFrozenPane::new(2, 0, "A2"), + }, + ]) + .unwrap_err(); + assert_eq!(error.code, "use.office.spreadsheet_freeze_geometry_invalid"); + assert_eq!(editor.package().content_sha256(), before); + assert!(editor.snapshot().unwrap().get("/Sheet1/A1", 0).is_err()); + + editor + .set_spreadsheet_frozen_pane("/Sheet1", NativeSpreadsheetFrozenPane::new(2, 0, "A3")) + .unwrap(); + let pane = editor.snapshot().unwrap().get("/Sheet1/freeze", 0).unwrap(); + assert_eq!(pane.format["activePane"], "bottomLeft"); + editor.remove("/Sheet1/freeze").unwrap(); + assert!(editor.snapshot().unwrap().get("/Sheet1/freeze", 0).is_err()); +} + +#[tokio::test] +async fn import_preserves_strict_spreadsheetml_and_freeze_fails_closed_on_unknown_content() { + const TRANSITIONAL: &str = "http://schemas.openxmlformats.org/spreadsheetml/2006/main"; + const STRICT: &str = "http://purl.oclc.org/ooxml/spreadsheetml/main"; + + let temp = tempfile::tempdir().unwrap(); + let path = temp.path().join("strict-import.xlsx"); + let editor = NativeOfficeEditor::create(&path).await.unwrap(); + let mut package = editor.package().clone(); + let worksheet = std::str::from_utf8(package.part("xl/worksheets/sheet1.xml").unwrap()) + .unwrap() + .replace(TRANSITIONAL, STRICT); + package + .set_part("xl/worksheets/sheet1.xml", worksheet.into_bytes()) + .unwrap(); + let mut editor = NativeOfficeEditor::from_package(package).unwrap(); + + editor + .import_spreadsheet_delimited( + "/Sheet1", + NativeSpreadsheetDelimitedImport::new( + "Date,Amount\n2026-07-17,42", + NativeSpreadsheetDelimitedFormat::Csv, + ) + .with_header(true), + ) + .unwrap(); + let worksheet = + std::str::from_utf8(editor.package().part("xl/worksheets/sheet1.xml").unwrap()).unwrap(); + assert!(worksheet.contains(STRICT)); + assert!(!worksheet.contains(TRANSITIONAL)); + assert!(worksheet.contains(" [--output ]`. Spill children +are read-only; replacing or removing an anchor clears its old spill. Failures +including cycles, unsupported or qualified functions, unsupported +structured-reference forms, external-workbook references, overlapping formula +storage, and invalid OOXML roll back the complete batch. Spreadsheet errors +such as `#DIV/0!` and `#SPILL!` remain typed cell values. The engine never +fetches external workbooks or invokes a scripting fallback. + +ListObject structured references resolve table names or display names. +`Table[Column]` selects one data column and `Table[[First]:[Last]]` selects a +contiguous data-column range. `#All`, `#Data`, `#Headers`, and `#Totals` select +structural rows. `Table[@Column]`, `Table[[#This Row],[Column]]`, and +table-local `[@Column]` select the current data row; table-local forms require +the formula cell to be inside the inferred table. Missing tables, columns, or +requested header/totals rows, disjoint columns, and non-canonical forms remain +fail-closed. + +Exact replay accepts canonical formula storage and canonical array anchors +only when their native cached results are present. It rejects physical storage +that typed mutations cannot reproduce byte-for-byte, including explicit +`t="normal"` formulas and uncached or malformed array anchors, with +`use.office.dump_unsupported`. + +Semantic Spreadsheet cell nodes expose `formulaCached` for formula anchors and +`valuePresent` for physical cached or inline values. This distinguishes an +explicit empty string from an absent cell value and lets issue analysis report +only formulas that truly lack an OOXML cache. + +Formula calculation is bounded to 8,192 formula characters, depth 128, and +8,192 AST nodes, with the same depth bound applied across nested named +references; 100,000 reference areas per value; 100,000 formula cells, +1,000,000 graph edges, and 1,000,000 formula-cell reference visits; 100,000 +materialized cells in one array or function call and 100,000 cumulative spill +children in one pass; 200,000 OOXML cell writes; and 1 MiB for one UTF-8 text +result. Formula text is additionally limited to 8 MiB cumulatively per +calculation pass. Cell set/remove accepts normalized A1 rectangular ranges of +at most 100,000 cells and rolls back the whole operation on error. Native Spreadsheet structure operations insert or delete at most 10,000 rows or columns, rename worksheets, reorder worksheets, and copy a worksheet after its source or at an explicit one-based position. Copy assigns new worksheet and @@ -218,6 +276,31 @@ multi-key mutation with persisted worksheet sort state. It is deliberately separate from worksheet/table filter-definition lifecycle and from formula calculation. +Native bounded CSV/TSV import is implemented through typed Rust, versioned +batch, dedicated CLI, and standard MCP surfaces. One import targets an existing +worksheet at an explicit A1 start cell, accepts at most 8 MiB of UTF-8 and a +100,000-cell rectangular extent, and parses BOM, CRLF, quoted delimiters, +embedded newlines, and doubled quotes without a scripting runtime. Unclosed, +misplaced, or trailing quote content fails the whole transaction instead of +being guessed. Ragged missing fields preserve cells beyond that source row; +explicit empty fields clear an existing target cell without materializing a new +blank cell. + +Import inference writes formulas, finite numbers, booleans, ISO dates/times, or +text as typed cell values. Dates honor the workbook's 1900/1904 date system and +receive a canonical date number format. Formulas pass the bounded cell-formula +parser, are stored, and are marked for recalculation. Import does not calculate +them implicitly; callers can append the native recalculation mutation to the +same atomic batch. Header mode atomically adds or replaces the worksheet +AutoFilter over the imported extent and installs one canonical frozen pane +below the header. Frozen panes have typed set/remove, semantic `/Sheet/freeze` +readback, selectors, versioned batch and MCP payloads, and exact replay. +Unsupported or vendor-extended pane content is readable with +`nativeMutable=false` and fails closed on mutation. Tests cover +strict/transitional SpreadsheetML, malformed input rollback, typed values, +explicit empty cells, filter/freeze lifecycle, exact replay, file/stdin CLI, +and an unsaved/save/close/reopen standard MCP session without OfficeCLI. + Native `replace-text` is implemented through one typed Rust, batch, CLI, and standard MCP contract. Literal mode performs case-sensitive, non-overlapping substring matching. Regex mode uses Rust's linear-time regular-expression @@ -344,7 +427,9 @@ SpreadsheetML, while cell-range and defined-name sources remain formulas. An optional leading `=` is removed from comparison and custom formulas. Valid ISO dates from 1900 through 9999 become serial dates using the workbook's declared 1900 or 1904 date system; 1900 mode retains Excel's historical leap-day offset. -`HH:MM` and `HH:MM:SS` values become day fractions. No formula is evaluated. +`HH:MM` and `HH:MM:SS` values become day fractions. Data-validation formula +predicates are stored but are not executed by this feature or by the cell +formula recalculation pass. One rule carries typed blank, input-message, error-message, error-style, and list-dropdown state. New A3S rules default `allowBlank`, `showInput`, @@ -376,7 +461,7 @@ This milestone is cell data-validation structure, not complete rich Spreadsheet or OfficeCLI parity. Table calculated columns and totals functions, unsupported imported sort-state variants and date-group/color/icon filter families, charts, pivot tables, slicers, -sparklines, formula evaluation, CSV/TSV import, and Excel layout fidelity remain +sparklines, data-validation predicate execution, and Excel layout fidelity remain separate work. Native `add-conditional-format` and `set-conditional-format` form a separate @@ -444,10 +529,11 @@ atomic rollback, exact canonical replay, CLI lifecycle and atomic batch, MCP schema conversion, and a complete standard MCP lifecycle with an unusable OfficeCLI provider path. -This milestone stores conditional-format formulas but does not evaluate them or -render Excel's visual result. It does not provide x14 advanced visual options, -table/chart/pivot formatting, formula calculation, or full Excel rendering and -layout fidelity, and therefore is not complete OfficeCLI or Spreadsheet parity. +This milestone stores conditional-format formulas but does not evaluate those +rule predicates or render Excel's visual result. It does not provide x14 +advanced visual options, table/chart/pivot formatting, or full Excel rendering +and layout fidelity, and therefore is not complete OfficeCLI or Spreadsheet +parity. Native `add-named-range` and `set-named-range` form a separate closed typed Rust, versioned batch, CLI, and standard MCP contract for Spreadsheet defined @@ -485,21 +571,22 @@ Excel namespace also rejects collisions with ListObject `name` or `displayName`. `_xlnm.*` print/filter definitions and `Slicer_*` sentinels are protected because their owning typed features must manage them. -Every mutation marks the workbook for full recalculation without evaluating a -formula. The loss-preserving writer keeps workbook child order, strict or -transitional SpreadsheetML QNames, untouched defined names, and unknown -attributes. Unknown collection children or non-text name content fail closed; -removing the final name also fails if deleting `definedNames` would discard -unknown collection attributes. Exact replay emits named ranges after worksheet -creation and reproduces the supported part map byte-for-byte. Tests cover +Every defined-name mutation marks the workbook for recalculation but does not +calculate by itself. An explicit native recalculation pass resolves supported +names referenced by cell formulas. The loss-preserving writer keeps workbook +child order, strict or transitional SpreadsheetML QNames, untouched defined +names, and unknown attributes. Unknown collection children or non-text name +content fail closed; removing the final name also fails if deleting +`definedNames` would discard unknown collection attributes. Exact replay emits +named ranges after worksheet creation and reproduces the supported part map +byte-for-byte. Tests cover scoped identity and ambiguity, validation and rollback, reserved names, unknown data, strict OOXML, table-name collisions, replay, native CLI batch atomicity, and a complete standard MCP unsaved/save/close lifecycle with an unusable OfficeCLI provider. -This milestone is defined-name lifecycle and storage, not a formula parser, -dependency graph, evaluator, external-link authoring, or complete rich -Spreadsheet parity. +This milestone is defined-name lifecycle and storage, not external-link +authoring or complete rich Spreadsheet parity. Native `add-spreadsheet-auto-filter` and `set-spreadsheet-auto-filter` form a closed typed Rust, versioned batch, CLI, and standard MCP contract. Ordinary @@ -617,6 +704,28 @@ metadata, custom styles, or final collection data that cannot be retained make the semantic node non-mutable or fail with `use.office.spreadsheet_table_unknown_content` instead of being flattened. +Replacement also preserves formula identity for common ListObject structured +references. It maps the old table `name` to the new `name`, the old effective +`displayName` to the new effective `displayName`, and old columns to new +columns by physical position. When the old aliases coincide, the new display +name is preferred. The editor applies those rewrites to worksheet cell +formulas, workbook defined names, conditional-format and data-validation +formulas, chart formulas, and calculated-column/totals-row formulas in table +parts. Formula string literals and external-workbook structured references are +not changed. + +Table-local forms such as `[@Qty]` are rewritten only for a formula cell inside +the old table range or another carrier whose ListObject ownership is provable. +An affected local reference with unknown ownership, or a local reference +across a table range/header/totals-row geometry change, fails atomically with +`use.office.spreadsheet_table_formula_rewrite_unsupported`. Geometry changes +with matching explicit references clear formula and chart caches before the +workbook is marked for full recalculation. Removal fails atomically with +`use.office.spreadsheet_table_referenced` while an explicit target or a +provably applicable/unknown local structured reference remains. A qualified +reference to another worksheet and a local reference provably owned by another +table do not block the target lifecycle. + Semantic reads expose table nodes and stable child column nodes, including name/display name, normalized range, header/totals state, filter children, built-in style, display flags, table ID, and `nativeMutable`. Query supports @@ -933,10 +1042,11 @@ Root-scoped replay dump is implemented for the canonical subset that current typed mutations can reproduce exactly: plain Word paragraphs and rectangular tables, Spreadsheet worksheets, typed defined names, typed cells, typed worksheet/table AutoFilters, ListObject tables, stable physical row order with -supported typed sort state, merged ranges, typed data-validation rules, and -canonical typed conditional-format rules without cached formula results, and -Presentation slides with plain one-run text -shapes and canonical basic tables. The versioned +supported typed sort state, canonical frozen panes and import date styles, +merged ranges, typed data-validation rules, and +canonical typed conditional-format rules, natively recalculable formula caches +and canonical cached dynamic-array spills, and Presentation slides with plain +one-run text shapes and canonical basic tables. The versioned artifact records document kind, `/` scope, blank-template part-map SHA-256, ordered mutations, and expected result part-map SHA-256. Native `batch` checks both fingerprints and restores the original package on a failed result check. @@ -969,10 +1079,12 @@ coverage above, advanced image mutation and OOXML SVG fallback, complex/custom part carriers, Presentation table merges/rich styles, subtree and rich-structure dump, advanced rich-format operations, modern threaded comments and legacy-comment replies/resolution/rich bodies, and -the formula parser/dependency/recalculation engine remain before their -respective gates can be promoted. Creation and structural mutation remain -under the interoperability gate until Microsoft Office and optional CI -LibreOffice checks confirm that no repair dialog is required. +complete Excel function breadth, structured-reference forms beyond common row +items and contiguous columns, qualified functions, and external-workbook +formula calculation remain before their respective gates can be promoted. +Creation and structural mutation remain under the interoperability gate until +Microsoft Office and optional CI LibreOffice checks confirm that no repair +dialog is required. ### Gate 3 — Rich Word @@ -1050,6 +1162,7 @@ add/set/remove/move/copy/swap, scoped cross-format literal/regex replacement, cross-format text formatting, typed Spreadsheet number/fill/border/alignment and cell-presentation formatting, exact Spreadsheet merged-cell editing, typed Spreadsheet physical row sorting with persisted sort state, Spreadsheet +CSV/TSV import with typed inference and header filter/freeze behavior, worksheet/table AutoFilters, data-validation, conditional-formatting, and scoped defined-name editing, typed ListObject table lifecycle, inert hyperlinks, diff --git a/src/mcp/office/input.rs b/src/mcp/office/input.rs index 1422b02f..544e4291 100644 --- a/src/mcp/office/input.rs +++ b/src/mcp/office/input.rs @@ -16,13 +16,17 @@ mod cell_format; mod conditional_formatting; mod data_validation; mod spreadsheet_filter; +mod spreadsheet_import; mod spreadsheet_sort; +mod spreadsheet_view; use cell_format::OfficeCellFormat; use conditional_formatting::OfficeConditionalFormat; use data_validation::OfficeDataValidation; use spreadsheet_filter::{OfficeSpreadsheetAutoFilter, OfficeSpreadsheetFilterColumn}; +use spreadsheet_import::OfficeSpreadsheetDelimitedImport; use spreadsheet_sort::OfficeSpreadsheetSort; +use spreadsheet_view::OfficeSpreadsheetFrozenPane; const MAX_IMAGE_BYTES: usize = 64 * 1024 * 1024; @@ -785,6 +789,8 @@ pub(super) enum OfficeMutation { path: String, value: OfficeCellValue, }, + /// Calculate supported formulas and write cached values and dynamic spills. + RecalculateSpreadsheetFormulas, AddSpreadsheetTable { /// Existing Spreadsheet worksheet path such as `/Sheet1`. sheet: String, @@ -810,6 +816,16 @@ pub(super) enum OfficeMutation { path: String, sort: OfficeSpreadsheetSort, }, + ImportSpreadsheetDelimited { + /// Existing Spreadsheet worksheet path such as `/Sheet1`. + sheet: String, + import: OfficeSpreadsheetDelimitedImport, + }, + SetSpreadsheetFrozenPane { + /// Existing Spreadsheet worksheet path such as `/Sheet1`. + sheet: String, + pane: OfficeSpreadsheetFrozenPane, + }, AddNamedRange { #[serde(rename = "namedRange")] named_range: OfficeNamedRange, @@ -986,6 +1002,9 @@ impl OfficeMutation { path, value: value.into(), }, + Self::RecalculateSpreadsheetFormulas => { + NativeOfficeMutation::RecalculateSpreadsheetFormulas + } Self::AddSpreadsheetTable { sheet, table } => { NativeOfficeMutation::AddSpreadsheetTable { sheet, @@ -1016,6 +1035,18 @@ impl OfficeMutation { sort: sort.into(), } } + Self::ImportSpreadsheetDelimited { sheet, import } => { + NativeOfficeMutation::ImportSpreadsheetDelimited { + sheet, + import: import.into_native(), + } + } + Self::SetSpreadsheetFrozenPane { sheet, pane } => { + NativeOfficeMutation::SetSpreadsheetFrozenPane { + sheet, + pane: pane.into_native(), + } + } Self::AddNamedRange { named_range } => NativeOfficeMutation::AddNamedRange { named_range: named_range.into_native()?, }, diff --git a/src/mcp/office/input/spreadsheet_import.rs b/src/mcp/office/input/spreadsheet_import.rs new file mode 100644 index 00000000..9337fa41 --- /dev/null +++ b/src/mcp/office/input/spreadsheet_import.rs @@ -0,0 +1,36 @@ +use a3s_use_office::{NativeSpreadsheetDelimitedFormat, NativeSpreadsheetDelimitedImport}; +use serde::{Deserialize, Serialize}; + +#[derive(Debug, Clone, Copy, Serialize, Deserialize, schemars::JsonSchema)] +#[serde(rename_all = "kebab-case")] +pub(in crate::mcp::office) enum OfficeSpreadsheetDelimitedFormat { + Csv, + Tsv, +} + +#[derive(Debug, Clone, Serialize, Deserialize, schemars::JsonSchema)] +#[serde(rename_all = "camelCase", deny_unknown_fields)] +pub(in crate::mcp::office) struct OfficeSpreadsheetDelimitedImport { + /// Bounded UTF-8 CSV or TSV content; filesystem paths belong at the CLI boundary. + content: String, + format: OfficeSpreadsheetDelimitedFormat, + /// Treat the first imported row as headers, add an AutoFilter, and freeze below it. + #[serde(default)] + header: bool, + /// A1 cell at which the first source field is written. Defaults to A1. + start_cell: Option, +} + +impl OfficeSpreadsheetDelimitedImport { + pub(super) fn into_native(self) -> NativeSpreadsheetDelimitedImport { + NativeSpreadsheetDelimitedImport::new( + self.content, + match self.format { + OfficeSpreadsheetDelimitedFormat::Csv => NativeSpreadsheetDelimitedFormat::Csv, + OfficeSpreadsheetDelimitedFormat::Tsv => NativeSpreadsheetDelimitedFormat::Tsv, + }, + ) + .with_header(self.header) + .with_start_cell(self.start_cell.unwrap_or_else(|| "A1".into())) + } +} diff --git a/src/mcp/office/input/spreadsheet_view.rs b/src/mcp/office/input/spreadsheet_view.rs new file mode 100644 index 00000000..d1d6c9e1 --- /dev/null +++ b/src/mcp/office/input/spreadsheet_view.rs @@ -0,0 +1,19 @@ +use a3s_use_office::NativeSpreadsheetFrozenPane; +use serde::{Deserialize, Serialize}; + +#[derive(Debug, Clone, Serialize, Deserialize, schemars::JsonSchema)] +#[serde(rename_all = "camelCase", deny_unknown_fields)] +pub(in crate::mcp::office) struct OfficeSpreadsheetFrozenPane { + /// Number of complete rows frozen above the scrollable pane. + frozen_rows: u32, + /// Number of complete columns frozen to the left of the scrollable pane. + frozen_columns: u32, + /// First visible cell in the scrollable pane. + top_left_cell: String, +} + +impl OfficeSpreadsheetFrozenPane { + pub(super) fn into_native(self) -> NativeSpreadsheetFrozenPane { + NativeSpreadsheetFrozenPane::new(self.frozen_rows, self.frozen_columns, self.top_left_cell) + } +} diff --git a/src/mcp/office/tests.rs b/src/mcp/office/tests.rs index 10827174..8531937b 100644 --- a/src/mcp/office/tests.rs +++ b/src/mcp/office/tests.rs @@ -339,6 +339,98 @@ fn office_batch_schema_exposes_typed_spreadsheet_sorting() { assert!(unknown.is_err()); } +#[test] +fn office_batch_schema_exposes_native_spreadsheet_recalculation() { + let schema = schemars::schema_for!(OfficeBatchInput); + let encoded = serde_json::to_string(&schema).unwrap(); + assert!( + encoded.contains("recalculate-spreadsheet-formulas"), + "{encoded}" + ); + + let input: OfficeBatchInput = serde_json::from_value(serde_json::json!({ + "session": "workbook", + "mutations": [{ + "operation": "recalculate-spreadsheet-formulas" + }] + })) + .unwrap(); + assert!(matches!( + input.mutations[0].clone().into_native().unwrap(), + NativeOfficeMutation::RecalculateSpreadsheetFormulas + )); +} + +#[test] +fn office_batch_schema_exposes_typed_spreadsheet_import_and_frozen_panes() { + let schema = schemars::schema_for!(OfficeBatchInput); + let encoded = serde_json::to_string(&schema).unwrap(); + for expected in [ + "import-spreadsheet-delimited", + "set-spreadsheet-frozen-pane", + "startCell", + "frozenRows", + "frozenColumns", + "topLeftCell", + "csv", + "tsv", + ] { + assert!(encoded.contains(expected), "missing {expected}"); + } + + let input: OfficeBatchInput = serde_json::from_value(serde_json::json!({ + "session": "workbook", + "mutations": [ + { + "operation": "import-spreadsheet-delimited", + "sheet": "/Sheet1", + "import": { + "content": "Name,Value\nAlpha,42", + "format": "csv", + "header": true, + "startCell": "B2" + } + }, + { + "operation": "set-spreadsheet-frozen-pane", + "sheet": "/Sheet1", + "pane": { + "frozenRows": 2, + "frozenColumns": 0, + "topLeftCell": "B3" + } + } + ] + })) + .unwrap(); + assert!(matches!( + input.mutations[0].clone().into_native().unwrap(), + NativeOfficeMutation::ImportSpreadsheetDelimited { ref sheet, ref import } + if sheet == "/Sheet1" + && import.header + && import.start_cell == "B2" + && import.format == a3s_use_office::NativeSpreadsheetDelimitedFormat::Csv + )); + assert!(matches!( + input.mutations[1].clone().into_native().unwrap(), + NativeOfficeMutation::SetSpreadsheetFrozenPane { ref sheet, ref pane } + if sheet == "/Sheet1" + && pane.frozen_rows == 2 + && pane.frozen_columns == 0 + && pane.top_left_cell == "B3" + )); + + let unknown = serde_json::from_value::(serde_json::json!({ + "session": "workbook", + "mutations": [{ + "operation": "import-spreadsheet-delimited", + "sheet": "/Sheet1", + "import": {"content": "a", "format": "json"} + }] + })); + assert!(unknown.is_err()); +} + #[test] fn office_batch_schema_exposes_typed_spreadsheet_data_validation() { let schema = schemars::schema_for!(OfficeBatchInput); diff --git a/src/office_native_cli.rs b/src/office_native_cli.rs index 3c6a6063..d202605f 100644 --- a/src/office_native_cli.rs +++ b/src/office_native_cli.rs @@ -20,6 +20,8 @@ mod part; mod raw; mod replay; mod spreadsheet_filter; +mod spreadsheet_formula; +mod spreadsheet_import; mod spreadsheet_sort; mod spreadsheet_table; mod view; @@ -58,6 +60,8 @@ const HELP: &str = concat!( " a3s-use office native set [--range ] [--filter ...|--clear-filters] [--output ] [--json]\n", " a3s-use office native set [--name ] [--display-name ] [--range ] [--table-column ...] [--filter ...|--clear-filters] [--header-row ] [--totals-row ] [--style none|light:<1-21>|medium:<1-28>|dark:<1-11>] [--show-first-column ] [--show-last-column ] [--show-row-stripes ] [--show-column-stripes ] [--output ] [--json]\n", " a3s-use office native sort --key [--key ...] [--header ] [--case-sensitive ] [--output ] [--json]\n", + " a3s-use office native import [source.csv|source.tsv] [--file ] [--stdin] [--format csv|tsv] [--header] [--start-cell ] [--output ] [--json]\n", + " a3s-use office native recalculate [--output ] [--json]\n", " a3s-use office native remove [--output ] [--json]\n", " a3s-use office native move [--to ] [--index |--before |--after ] [--output ] [--json]\n", " a3s-use office native copy [--to ] [--name ] [--index |--before |--after ] [--output ] [--json]\n", @@ -87,6 +91,8 @@ pub async fn run(args: &[String]) -> UseResult { Some("add-part") => part::add(args).await, Some("set") => set(args).await, Some("sort") => spreadsheet_sort::run(args).await, + Some("import") => spreadsheet_import::run(args).await, + Some("recalculate") => spreadsheet_formula::recalculate(args).await, Some("remove") => remove(args).await, Some("move") => arrange::move_node(args).await, Some("copy") => arrange::copy_node(args).await, @@ -110,7 +116,7 @@ fn help() -> CommandOutput { HELP, serde_json::json!({ "commands": [ - "get", "query", "view", "watch", "raw", "raw-set", "dump", "merge", "validate", "create", "add", "add-part", "set", "sort", "remove", "move", "copy", "swap", + "get", "query", "view", "watch", "raw", "raw-set", "dump", "merge", "validate", "create", "add", "add-part", "set", "sort", "import", "recalculate", "remove", "move", "copy", "swap", "insert-rows", "delete-rows", "insert-columns", "delete-columns", "rename-sheet", "move-sheet", "copy-sheet", "batch" ], diff --git a/src/office_native_cli/bounded_input.rs b/src/office_native_cli/bounded_input.rs index 033ef493..f219fb24 100644 --- a/src/office_native_cli/bounded_input.rs +++ b/src/office_native_cli/bounded_input.rs @@ -6,6 +6,7 @@ pub(super) enum NativeInputKind { Batch, Image, RawXml, + SpreadsheetImport, TemplateData, } @@ -15,6 +16,7 @@ impl NativeInputKind { Self::Batch => "use.office.batch_input", Self::Image => "use.office.image_input", Self::RawXml => "use.office.raw_input", + Self::SpreadsheetImport => "use.office.spreadsheet_import_input", Self::TemplateData => "use.office.template_data_input", } } @@ -24,11 +26,29 @@ impl NativeInputKind { Self::Batch => "Native Office batch input", Self::Image => "Native Office image input", Self::RawXml => "Native Office raw XML input", + Self::SpreadsheetImport => "Native Spreadsheet delimited import input", Self::TemplateData => "Native Office template data input", } } } +pub(super) async fn read_bounded_stdin(limit: u64, kind: NativeInputKind) -> UseResult> { + let mut bytes = Vec::new(); + let mut reader = tokio::io::stdin().take(limit + 1); + reader.read_to_end(&mut bytes).await.map_err(|error| { + input_error( + kind, + "read_failed", + "", + format!("Failed to read {} from stdin: {error}", kind.label()), + ) + })?; + if bytes.len() as u64 > limit { + return Err(input_too_large(kind, "", limit)); + } + Ok(bytes) +} + pub(super) async fn read_bounded_input( path: &str, limit: u64, diff --git a/src/office_native_cli/spreadsheet_formula.rs b/src/office_native_cli/spreadsheet_formula.rs new file mode 100644 index 00000000..01a3ef8d --- /dev/null +++ b/src/office_native_cli/spreadsheet_formula.rs @@ -0,0 +1,40 @@ +use a3s_use_core::UseResult; +use a3s_use_office::NativeOfficeEditor; + +use super::arguments::{AllowedOptions, ParsedArguments}; +use super::{save_editor, usage_error}; +use crate::cli::CommandOutput; + +pub(super) async fn recalculate(args: &[String]) -> UseResult { + let parsed = ParsedArguments::parse(args, AllowedOptions::MUTATE)?; + if parsed.positionals.len() != 1 { + return Err(usage_error( + "office native recalculate requires ", + )); + } + let source = &parsed.positionals[0]; + let mut editor = NativeOfficeEditor::open(source).await?; + let source_path = editor.package().path().to_path_buf(); + let calculation = editor.recalculate_spreadsheet_formulas()?; + let changed = editor.is_dirty(); + save_editor(&mut editor, parsed.output.as_deref()).await?; + let output_path = editor.package().path().to_path_buf(); + let in_place = output_path == source_path; + Ok(CommandOutput::success( + format!( + "Recalculated {} formula(s) and {} spill cell(s), then saved '{}'.", + calculation.formula_count, + calculation.spill_cell_count, + output_path.display() + ), + serde_json::json!({ + "operation": "recalculate-spreadsheet-formulas", + "changed": changed, + "result": calculation, + "kind": editor.package().kind(), + "outputPath": output_path, + "inPlace": in_place, + "revision": editor.package().source_revision() + }), + )) +} diff --git a/src/office_native_cli/spreadsheet_import.rs b/src/office_native_cli/spreadsheet_import.rs new file mode 100644 index 00000000..9d2abd3a --- /dev/null +++ b/src/office_native_cli/spreadsheet_import.rs @@ -0,0 +1,251 @@ +use std::path::Path; + +use a3s_use_core::UseResult; +use a3s_use_office::{ + NativeOfficeEditor, NativeSpreadsheetDelimitedFormat, NativeSpreadsheetDelimitedImport, + MAX_NATIVE_SPREADSHEET_IMPORT_BYTES, +}; + +use super::bounded_input::{input_error, read_bounded_input, read_bounded_stdin, NativeInputKind}; +use super::{save_editor, usage_error}; +use crate::cli::CommandOutput; + +#[derive(Debug, Default, PartialEq, Eq)] +struct ImportArguments { + positionals: Vec, + source_file: Option, + stdin: bool, + format: Option, + header: bool, + start_cell: String, + output: Option, +} + +pub(super) async fn run(args: &[String]) -> UseResult { + let parsed = parse(args)?; + if !(2..=3).contains(&parsed.positionals.len()) { + return Err(usage_error( + "office native import requires , , and one source file or --stdin", + )); + } + let positional_source = parsed.positionals.get(2).map(String::as_str); + let sources = usize::from(positional_source.is_some()) + + usize::from(parsed.source_file.is_some()) + + usize::from(parsed.stdin); + if sources != 1 { + return Err(usage_error( + "office native import requires exactly one positional source file, --file , or --stdin", + )); + } + let source_file = parsed.source_file.as_deref().or(positional_source); + let limit = u64::try_from(MAX_NATIVE_SPREADSHEET_IMPORT_BYTES).unwrap_or(u64::MAX); + let bytes = if parsed.stdin { + read_bounded_stdin(limit, NativeInputKind::SpreadsheetImport).await? + } else { + read_bounded_input( + source_file.ok_or_else(|| usage_error("Spreadsheet import source is missing"))?, + limit, + NativeInputKind::SpreadsheetImport, + ) + .await? + }; + let input_label = source_file.unwrap_or(""); + let content = String::from_utf8(bytes).map_err(|error| { + input_error( + NativeInputKind::SpreadsheetImport, + "invalid", + input_label, + format!("Spreadsheet import input '{input_label}' is not valid UTF-8: {error}"), + ) + })?; + let format = parsed.format.unwrap_or_else(|| infer_format(source_file)); + let import = NativeSpreadsheetDelimitedImport::new(content, format) + .with_header(parsed.header) + .with_start_cell(parsed.start_cell); + + let source = &parsed.positionals[0]; + let sheet = &parsed.positionals[1]; + let mut editor = NativeOfficeEditor::open(source).await?; + let source_path = editor.package().path().to_path_buf(); + let result = editor.import_spreadsheet_delimited(sheet, import)?; + save_editor(&mut editor, parsed.output.as_deref()).await?; + let output_path = editor.package().path().to_path_buf(); + let in_place = output_path == source_path; + let human = if result.changed { + format!( + "Imported {} row(s) x {} column(s) into {} at {} and saved '{}'.", + result.row_count, + result.column_count, + result.sheet, + result.start_cell, + output_path.display() + ) + } else { + format!( + "No delimited rows were present; saved '{}'.", + output_path.display() + ) + }; + Ok(CommandOutput::success( + human, + serde_json::json!({ + "operation": "import-spreadsheet-delimited", + "changed": result.changed, + "source": if parsed.stdin { serde_json::Value::String("stdin".into()) } else { serde_json::Value::String(input_label.into()) }, + "result": result, + "kind": editor.package().kind(), + "outputPath": output_path, + "inPlace": in_place, + "revision": editor.package().source_revision() + }), + )) +} + +fn parse(args: &[String]) -> UseResult { + let mut parsed = ImportArguments { + start_cell: "A1".into(), + ..ImportArguments::default() + }; + let mut index = 1; + let mut header = false; + let mut start_cell_seen = false; + while index < args.len() { + match args[index].as_str() { + "--json" => index += 1, + "--" => { + parsed.positionals.extend_from_slice(&args[index + 1..]); + break; + } + "--file" => { + set_option(&mut parsed.source_file, args, index, "--file")?; + index += 2; + } + "--stdin" => { + if parsed.stdin { + return Err(usage_error("--stdin may be specified only once")); + } + parsed.stdin = true; + index += 1; + } + "--format" => { + if parsed.format.is_some() { + return Err(usage_error("--format may be specified only once")); + } + let value = option_value(args, index, "--format")?; + parsed.format = Some(match value.to_ascii_lowercase().as_str() { + "csv" => NativeSpreadsheetDelimitedFormat::Csv, + "tsv" => NativeSpreadsheetDelimitedFormat::Tsv, + _ => return Err(usage_error("--format requires csv or tsv")), + }); + index += 2; + } + "--header" => { + if header { + return Err(usage_error("--header may be specified only once")); + } + header = true; + parsed.header = true; + index += 1; + } + "--start-cell" => { + if start_cell_seen { + return Err(usage_error("--start-cell may be specified only once")); + } + parsed.start_cell = option_value(args, index, "--start-cell")?.into(); + start_cell_seen = true; + index += 2; + } + "--output" => { + set_option(&mut parsed.output, args, index, "--output")?; + index += 2; + } + option if option.starts_with('-') => { + return Err(usage_error(format!( + "unknown native Spreadsheet import option '{option}'" + ))); + } + positional => { + parsed.positionals.push(positional.into()); + index += 1; + } + } + } + Ok(parsed) +} + +fn infer_format(source: Option<&str>) -> NativeSpreadsheetDelimitedFormat { + let extension = source + .and_then(|source| Path::new(source).extension()) + .and_then(|extension| extension.to_str()) + .map(str::to_ascii_lowercase); + if matches!(extension.as_deref(), Some("tsv" | "tab")) { + NativeSpreadsheetDelimitedFormat::Tsv + } else { + NativeSpreadsheetDelimitedFormat::Csv + } +} + +fn set_option( + target: &mut Option, + args: &[String], + index: usize, + option: &str, +) -> UseResult<()> { + if target.is_some() { + return Err(usage_error(format!("{option} may be specified only once"))); + } + *target = Some(option_value(args, index, option)?.into()); + Ok(()) +} + +fn option_value<'a>(args: &'a [String], index: usize, option: &str) -> UseResult<&'a str> { + args.get(index + 1) + .filter(|value| !value.starts_with("--")) + .map(String::as_str) + .ok_or_else(|| usage_error(format!("{option} requires a value"))) +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn parses_file_stdin_format_header_and_start_cell() { + let parsed = parse(&[ + "import".into(), + "book.xlsx".into(), + "/Data".into(), + "--file".into(), + "source.tab".into(), + "--format".into(), + "tsv".into(), + "--header".into(), + "--start-cell".into(), + "B2".into(), + "--output".into(), + "copy.xlsx".into(), + ]) + .unwrap(); + assert_eq!(parsed.positionals, ["book.xlsx", "/Data"]); + assert_eq!(parsed.source_file.as_deref(), Some("source.tab")); + assert_eq!(parsed.format, Some(NativeSpreadsheetDelimitedFormat::Tsv)); + assert!(parsed.header); + assert_eq!(parsed.start_cell, "B2"); + assert_eq!(parsed.output.as_deref(), Some("copy.xlsx")); + assert_eq!( + infer_format(Some("data.TAB")), + NativeSpreadsheetDelimitedFormat::Tsv + ); + } + + #[test] + fn rejects_duplicate_and_unknown_options() { + for args in [ + vec!["import".into(), "--stdin".into(), "--stdin".into()], + vec!["import".into(), "--format".into(), "json".into()], + vec!["import".into(), "--unknown".into()], + ] { + assert_eq!(parse(&args).unwrap_err().code, "use.cli.invalid_usage"); + } + } +} diff --git a/tests/cli.rs b/tests/cli.rs index 3bdb785f..a35de3f5 100644 --- a/tests/cli.rs +++ b/tests/cli.rs @@ -1126,6 +1126,31 @@ fn native_office_cli_writes_typed_spreadsheet_values_without_an_officecli_provid let formula: serde_json::Value = serde_json::from_slice(&formula.stdout).unwrap(); assert_eq!(formula["data"]["node"]["format"]["formula"], "A1*2"); + let before_invalid_formula = std::fs::read(&document).unwrap(); + let invalid_formula = Command::new(binary()) + .args([ + "office", + "native", + "set", + document.to_str().unwrap(), + "/Sheet1/C2", + "--formula", + "SUM(A1", + "--json", + ]) + .env("A3S_OFFICECLI_EXECUTABLE", &provider) + .output() + .unwrap(); + assert!(!invalid_formula.status.success(), "{invalid_formula:?}"); + let invalid_formula: serde_json::Value = + serde_json::from_slice(&invalid_formula.stdout).unwrap(); + assert_eq!( + invalid_formula["error"]["code"], + "use.office.spreadsheet_formula_invalid" + ); + assert_eq!(invalid_formula["error"]["details"]["characterOffset"], 6); + assert_eq!(std::fs::read(&document).unwrap(), before_invalid_formula); + std::fs::write( &mutations, serde_json::to_vec(&serde_json::json!({ diff --git a/tests/office_cell_format_mcp.rs b/tests/office_cell_format_mcp.rs index f154c569..cd3fb2fd 100644 --- a/tests/office_cell_format_mcp.rs +++ b/tests/office_cell_format_mcp.rs @@ -157,6 +157,42 @@ async fn native_standard_mcp_applies_typed_spreadsheet_cell_format_without_offic ); assert_eq!(applied["result"]["structuredContent"]["persisted"], false); + let invalid_formula = call( + &mut stdin, + &mut stdout, + 40, + "office_apply_batch", + serde_json::json!({ + "session": "workbook", + "mutations": [ + { + "operation": "set-cell-value", + "path": "/Sheet1/A1", + "value": { "type": "number", "value": "999" } + }, + { + "operation": "set-cell-value", + "path": "/Sheet1/B1", + "value": { "type": "formula", "expression": "SUM(A1" } + } + ] + }), + TIMEOUT, + ) + .await; + assert_eq!( + invalid_formula["result"]["isError"], true, + "{invalid_formula}" + ); + assert_eq!( + invalid_formula["result"]["structuredContent"]["code"], + "use.office.spreadsheet_formula_invalid" + ); + assert_eq!( + invalid_formula["result"]["structuredContent"]["details"]["characterOffset"], + 6 + ); + let read = call( &mut stdin, &mut stdout, diff --git a/tests/office_spreadsheet_formula_cli.rs b/tests/office_spreadsheet_formula_cli.rs new file mode 100644 index 00000000..0e7cce4c --- /dev/null +++ b/tests/office_spreadsheet_formula_cli.rs @@ -0,0 +1,159 @@ +#![cfg(feature = "office")] + +use std::path::Path; +use std::process::{Command, Output}; + +fn binary() -> &'static str { + env!("CARGO_BIN_EXE_a3s-use") +} + +fn execute(provider: &Path, args: &[&str]) -> Output { + Command::new(binary()) + .args(args) + .env("A3S_OFFICECLI_EXECUTABLE", provider) + .output() + .unwrap() +} + +fn success(output: Output) -> serde_json::Value { + assert!(output.status.success(), "{output:?}"); + serde_json::from_slice(&output.stdout).unwrap() +} + +fn run(provider: &Path, args: &[&str]) -> serde_json::Value { + success(execute(provider, args)) +} + +#[test] +fn native_cli_recalculates_typed_formulas_and_spills_without_officecli() { + let temp = tempfile::tempdir().unwrap(); + let provider = temp.path().join("must-not-be-invoked"); + let source = temp.path().join("formulas.xlsx"); + let output = temp.path().join("calculated.xlsx"); + run( + &provider, + &[ + "office", + "native", + "create", + source.to_str().unwrap(), + "--json", + ], + ); + for (path, option, value) in [ + ("/Sheet1/A1", "--number", "2"), + ("/Sheet1/B1", "--formula", "A1*3"), + ("/Sheet1/C1", "--formula", "SEQUENCE(2,2,1,1)"), + ] { + run( + &provider, + &[ + "office", + "native", + "set", + source.to_str().unwrap(), + path, + option, + value, + "--json", + ], + ); + } + + let recalculated = run( + &provider, + &[ + "office", + "native", + "recalculate", + source.to_str().unwrap(), + "--output", + output.to_str().unwrap(), + "--json", + ], + ); + assert_eq!( + recalculated["data"]["operation"], + "recalculate-spreadsheet-formulas" + ); + assert_eq!(recalculated["data"]["result"]["formulaCount"], 2); + assert_eq!(recalculated["data"]["result"]["spillCellCount"], 3); + assert_eq!(recalculated["data"]["inPlace"], false); + + let b1 = run( + &provider, + &[ + "office", + "native", + "get", + output.to_str().unwrap(), + "/Sheet1/B1", + "--json", + ], + ); + assert_eq!(b1["data"]["node"]["text"], "6"); + assert_eq!(b1["data"]["node"]["format"]["formulaCached"], "true"); + let d2 = run( + &provider, + &[ + "office", + "native", + "get", + output.to_str().unwrap(), + "/Sheet1/D2", + "--json", + ], + ); + assert_eq!(d2["data"]["node"]["text"], "4"); + assert!(d2["data"]["node"]["format"].get("formula").is_none()); + assert!(!provider.exists()); +} + +#[test] +fn native_cli_recalculation_failure_does_not_change_the_file() { + let temp = tempfile::tempdir().unwrap(); + let provider = temp.path().join("must-not-be-invoked"); + let document = temp.path().join("unsupported.xlsx"); + run( + &provider, + &[ + "office", + "native", + "create", + document.to_str().unwrap(), + "--json", + ], + ); + run( + &provider, + &[ + "office", + "native", + "set", + document.to_str().unwrap(), + "/Sheet1/A1", + "--formula", + "SHELL(\"unsafe\")", + "--json", + ], + ); + let before = std::fs::read(&document).unwrap(); + let failure = execute( + &provider, + &[ + "office", + "native", + "recalculate", + document.to_str().unwrap(), + "--json", + ], + ); + assert!(!failure.status.success(), "{failure:?}"); + let error: serde_json::Value = serde_json::from_slice(&failure.stdout).unwrap(); + assert_eq!( + error["error"]["code"], + "use.office.spreadsheet_formula_function_unsupported" + ); + assert_eq!(std::fs::read(&document).unwrap(), before); + assert!(!provider.exists()); +} diff --git a/tests/office_spreadsheet_formula_mcp.rs b/tests/office_spreadsheet_formula_mcp.rs new file mode 100644 index 00000000..885fd4b9 --- /dev/null +++ b/tests/office_spreadsheet_formula_mcp.rs @@ -0,0 +1,249 @@ +#![cfg(all(feature = "office", feature = "mcp"))] + +use std::process::Stdio; +use std::time::Duration; + +use tokio::io::{AsyncBufReadExt, AsyncReadExt, AsyncWriteExt, BufReader}; + +fn binary() -> &'static str { + env!("CARGO_BIN_EXE_a3s-use") +} + +#[tokio::test] +async fn standard_mcp_recalculates_formulas_atomically_without_officecli() { + const TIMEOUT: Duration = Duration::from_secs(15); + + let temp = tempfile::tempdir().unwrap(); + let provider = temp.path().join("must-not-be-invoked"); + let document = temp.path().join("mcp-formulas.xlsx"); + let mut child = tokio::process::Command::new(binary()) + .args(["mcp", "serve", "office-native"]) + .env("A3S_OFFICECLI_EXECUTABLE", &provider) + .stdin(Stdio::piped()) + .stdout(Stdio::piped()) + .stderr(Stdio::piped()) + .kill_on_drop(true) + .spawn() + .unwrap(); + let mut stdin = child.stdin.take().unwrap(); + let mut stdout = BufReader::new(child.stdout.take().unwrap()); + let mut stderr = child.stderr.take().unwrap(); + + request( + &mut stdin, + &mut stdout, + serde_json::json!({ + "jsonrpc": "2.0", + "id": 1, + "method": "initialize", + "params": { + "protocolVersion": "2025-06-18", + "capabilities": {}, + "clientInfo": { "name": "office-formula-test", "version": "1" } + } + }), + TIMEOUT, + ) + .await; + stdin + .write_all( + b"{\"jsonrpc\":\"2.0\",\"method\":\"notifications/initialized\",\"params\":{}}\n", + ) + .await + .unwrap(); + stdin.flush().await.unwrap(); + + let tools = request( + &mut stdin, + &mut stdout, + serde_json::json!({"jsonrpc":"2.0","id":2,"method":"tools/list","params":{}}), + TIMEOUT, + ) + .await; + let schema = tools["result"]["tools"] + .as_array() + .unwrap() + .iter() + .find(|tool| tool["name"] == "office_apply_batch") + .unwrap()["inputSchema"] + .to_string(); + assert!(schema.contains("recalculate-spreadsheet-formulas")); + + let created = call( + &mut stdin, + &mut stdout, + 3, + "office_create", + serde_json::json!({"session":"workbook","file":document}), + TIMEOUT, + ) + .await; + assert_ne!(created["result"]["isError"], true, "{created}"); + + let applied = call( + &mut stdin, + &mut stdout, + 4, + "office_apply_batch", + serde_json::json!({ + "session": "workbook", + "mutations": [ + { + "operation": "set-cell-value", + "path": "/Sheet1/A1", + "value": { "type": "number", "value": "2" } + }, + { + "operation": "set-cell-value", + "path": "/Sheet1/B1", + "value": { "type": "formula", "expression": "A1*3" } + }, + { + "operation": "set-cell-value", + "path": "/Sheet1/C1", + "value": { "type": "formula", "expression": "SEQUENCE(2,2,1,1)" } + }, + { + "operation": "recalculate-spreadsheet-formulas" + } + ] + }), + TIMEOUT, + ) + .await; + assert_ne!(applied["result"]["isError"], true, "{applied}"); + let content = &applied["result"]["structuredContent"]; + assert_eq!(content["persisted"], false); + assert_eq!(content["result"]["applied"], 4); + assert_eq!( + content["result"]["spreadsheetCalculations"][0]["formulaCount"], + 2 + ); + assert_eq!( + content["result"]["spreadsheetCalculations"][0]["spillCellCount"], + 3 + ); + + let spill = call( + &mut stdin, + &mut stdout, + 5, + "office_get", + serde_json::json!({"session":"workbook","path":"/Sheet1/D2","depth":0}), + TIMEOUT, + ) + .await; + assert_eq!(spill["result"]["structuredContent"]["node"]["text"], "4"); + + let rejected = call( + &mut stdin, + &mut stdout, + 6, + "office_apply_batch", + serde_json::json!({ + "session": "workbook", + "mutations": [ + { + "operation": "set-cell-value", + "path": "/Sheet1/E1", + "value": { "type": "text", "value": "must roll back" } + }, + { + "operation": "set-cell-value", + "path": "/Sheet1/F1", + "value": { "type": "formula", "expression": "SHELL(\"unsafe\")" } + }, + { + "operation": "recalculate-spreadsheet-formulas" + } + ] + }), + TIMEOUT, + ) + .await; + assert_eq!(rejected["result"]["isError"], true, "{rejected}"); + assert_eq!( + rejected["result"]["structuredContent"]["code"], + "use.office.spreadsheet_formula_function_unsupported" + ); + let rolled_back = call( + &mut stdin, + &mut stdout, + 7, + "office_get", + serde_json::json!({"session":"workbook","path":"/Sheet1/E1","depth":0}), + TIMEOUT, + ) + .await; + assert_eq!(rolled_back["result"]["isError"], true, "{rolled_back}"); + assert_eq!( + rolled_back["result"]["structuredContent"]["code"], + "use.office.node_not_found" + ); + + call( + &mut stdin, + &mut stdout, + 8, + "office_close", + serde_json::json!({"session":"workbook","discard":true}), + TIMEOUT, + ) + .await; + drop(stdin); + let status = tokio::time::timeout(TIMEOUT, child.wait()) + .await + .unwrap() + .unwrap(); + assert!(status.success()); + let mut diagnostics = Vec::new(); + stderr.read_to_end(&mut diagnostics).await.unwrap(); + assert!( + diagnostics.is_empty(), + "{}", + String::from_utf8_lossy(&diagnostics) + ); + assert!(!provider.exists()); +} + +async fn call( + stdin: &mut tokio::process::ChildStdin, + stdout: &mut BufReader, + id: u32, + name: &str, + arguments: serde_json::Value, + timeout: Duration, +) -> serde_json::Value { + request( + stdin, + stdout, + serde_json::json!({ + "jsonrpc": "2.0", + "id": id, + "method": "tools/call", + "params": {"name": name, "arguments": arguments} + }), + timeout, + ) + .await +} + +async fn request( + stdin: &mut tokio::process::ChildStdin, + stdout: &mut BufReader, + value: serde_json::Value, + timeout: Duration, +) -> serde_json::Value { + stdin + .write_all(format!("{value}\n").as_bytes()) + .await + .unwrap(); + stdin.flush().await.unwrap(); + let mut line = String::new(); + tokio::time::timeout(timeout, stdout.read_line(&mut line)) + .await + .unwrap() + .unwrap(); + assert!(!line.is_empty()); + serde_json::from_str(&line).unwrap() +} diff --git a/tests/office_spreadsheet_import_cli.rs b/tests/office_spreadsheet_import_cli.rs new file mode 100644 index 00000000..d9fbce50 --- /dev/null +++ b/tests/office_spreadsheet_import_cli.rs @@ -0,0 +1,275 @@ +#![cfg(feature = "office")] + +use std::io::Write; +use std::path::Path; +use std::process::{Command, Output, Stdio}; + +fn binary() -> &'static str { + env!("CARGO_BIN_EXE_a3s-use") +} + +fn execute(provider: &Path, args: &[&str]) -> Output { + Command::new(binary()) + .args(args) + .env("A3S_OFFICECLI_EXECUTABLE", provider) + .output() + .unwrap() +} + +fn execute_stdin(provider: &Path, args: &[&str], input: &[u8]) -> Output { + let mut child = Command::new(binary()) + .args(args) + .env("A3S_OFFICECLI_EXECUTABLE", provider) + .stdin(Stdio::piped()) + .stdout(Stdio::piped()) + .stderr(Stdio::piped()) + .spawn() + .unwrap(); + child.stdin.take().unwrap().write_all(input).unwrap(); + child.wait_with_output().unwrap() +} + +fn success(output: Output) -> serde_json::Value { + assert!(output.status.success(), "{output:?}"); + serde_json::from_slice(&output.stdout).unwrap() +} + +fn failure(output: Output) -> serde_json::Value { + assert!(!output.status.success(), "{output:?}"); + serde_json::from_slice(&output.stdout).unwrap() +} + +fn run(provider: &Path, args: &[&str]) -> serde_json::Value { + success(execute(provider, args)) +} + +#[test] +fn native_cli_imports_files_and_stdin_without_invoking_the_compatibility_provider() { + let temp = tempfile::tempdir().unwrap(); + let provider = temp.path().join("must-not-be-invoked"); + let document = temp.path().join("import.xlsx"); + let source = temp.path().join("source.tab"); + std::fs::write( + &source, + b"Name\tAmount\tDate\nAlpha\t42\t2026-07-17\nBeta\tTRUE\t2026-07-18", + ) + .unwrap(); + run( + &provider, + &[ + "office", + "native", + "create", + document.to_str().unwrap(), + "--json", + ], + ); + + let imported = run( + &provider, + &[ + "office", + "native", + "import", + document.to_str().unwrap(), + "/Sheet1", + source.to_str().unwrap(), + "--header", + "--start-cell", + "B2", + "--json", + ], + ); + assert_eq!( + imported["data"]["operation"], + "import-spreadsheet-delimited" + ); + assert_eq!(imported["data"]["result"]["format"], "tsv"); + assert_eq!(imported["data"]["result"]["range"], "B2:D4"); + assert_eq!(imported["data"]["result"]["rowCount"], 3); + assert_eq!(imported["data"]["result"]["freezePath"], "/Sheet1/freeze"); + + let amount = run( + &provider, + &[ + "office", + "native", + "get", + document.to_str().unwrap(), + "/Sheet1/C3", + "--json", + ], + ); + assert_eq!(amount["data"]["node"]["text"], "42"); + assert_eq!(amount["data"]["node"]["format"]["valueType"], "Number"); + let freeze = run( + &provider, + &[ + "office", + "native", + "get", + document.to_str().unwrap(), + "/Sheet1/freeze", + "--json", + ], + ); + assert_eq!(freeze["data"]["node"]["format"]["frozenRows"], "2"); + assert_eq!(freeze["data"]["node"]["format"]["topLeftCell"], "B3"); + let dump = run( + &provider, + &[ + "office", + "native", + "dump", + document.to_str().unwrap(), + "--json", + ], + ); + assert!(dump["data"]["artifact"]["mutations"] + .as_array() + .unwrap() + .iter() + .any(|mutation| mutation["operation"] == "set-spreadsheet-frozen-pane")); + + let stdin_document = temp.path().join("stdin.xlsx"); + let copy = temp.path().join("stdin-copy.xlsx"); + run( + &provider, + &[ + "office", + "native", + "create", + stdin_document.to_str().unwrap(), + "--json", + ], + ); + let output = execute_stdin( + &provider, + &[ + "office", + "native", + "import", + stdin_document.to_str().unwrap(), + "/Sheet1", + "--stdin", + "--format", + "csv", + "--output", + copy.to_str().unwrap(), + "--json", + ], + b"one,two\nthree,four", + ); + let imported = success(output); + assert_eq!(imported["data"]["source"], "stdin"); + assert_eq!(imported["data"]["inPlace"], false); + assert!(copy.is_file()); + let copied = run( + &provider, + &[ + "office", + "native", + "get", + copy.to_str().unwrap(), + "/Sheet1/B2", + "--json", + ], + ); + assert_eq!(copied["data"]["node"]["text"], "four"); + let original = failure(execute( + &provider, + &[ + "office", + "native", + "get", + stdin_document.to_str().unwrap(), + "/Sheet1/A1", + "--json", + ], + )); + assert_eq!(original["error"]["code"], "use.office.node_not_found"); + assert!(!provider.exists()); +} + +#[test] +fn native_cli_rejects_ambiguous_and_non_utf8_import_sources_before_mutation() { + let temp = tempfile::tempdir().unwrap(); + let provider = temp.path().join("must-not-be-invoked"); + let document = temp.path().join("safe.xlsx"); + let first = temp.path().join("first.csv"); + let second = temp.path().join("second.csv"); + let invalid = temp.path().join("invalid.csv"); + let malformed = temp.path().join("malformed.csv"); + std::fs::write(&first, b"a").unwrap(); + std::fs::write(&second, b"b").unwrap(); + std::fs::write(&invalid, [0xff, 0xfe]).unwrap(); + std::fs::write(&malformed, b"\"unterminated").unwrap(); + run( + &provider, + &[ + "office", + "native", + "create", + document.to_str().unwrap(), + "--json", + ], + ); + let before = std::fs::read(&document).unwrap(); + + let ambiguous = failure(execute( + &provider, + &[ + "office", + "native", + "import", + document.to_str().unwrap(), + "/Sheet1", + first.to_str().unwrap(), + "--file", + second.to_str().unwrap(), + "--json", + ], + )); + assert_eq!(ambiguous["error"]["code"], "use.cli.invalid_usage"); + assert_eq!(std::fs::read(&document).unwrap(), before); + + let invalid_utf8 = failure(execute( + &provider, + &[ + "office", + "native", + "import", + document.to_str().unwrap(), + "/Sheet1", + "--file", + invalid.to_str().unwrap(), + "--json", + ], + )); + assert_eq!( + invalid_utf8["error"]["code"], + "use.office.spreadsheet_import_input_invalid" + ); + assert_eq!(std::fs::read(&document).unwrap(), before); + + let malformed_csv = failure(execute( + &provider, + &[ + "office", + "native", + "import", + document.to_str().unwrap(), + "/Sheet1", + malformed.to_str().unwrap(), + "--json", + ], + )); + assert_eq!( + malformed_csv["error"]["code"], + "use.office.spreadsheet_import_delimited_invalid" + ); + assert_eq!(malformed_csv["error"]["details"]["row"], 1); + assert_eq!(malformed_csv["error"]["details"]["column"], 1); + assert_eq!(std::fs::read(&document).unwrap(), before); + assert!(!provider.exists()); +} diff --git a/tests/office_spreadsheet_import_mcp.rs b/tests/office_spreadsheet_import_mcp.rs new file mode 100644 index 00000000..e8a9abcb --- /dev/null +++ b/tests/office_spreadsheet_import_mcp.rs @@ -0,0 +1,331 @@ +#![cfg(all(feature = "office", feature = "mcp"))] + +use std::process::Stdio; +use std::time::Duration; + +use tokio::io::{AsyncBufReadExt, AsyncReadExt, AsyncWriteExt, BufReader}; + +fn binary() -> &'static str { + env!("CARGO_BIN_EXE_a3s-use") +} + +#[tokio::test] +async fn standard_mcp_imports_delimited_content_before_save_and_after_reopen() { + const TIMEOUT: Duration = Duration::from_secs(15); + + let temp = tempfile::tempdir().unwrap(); + let provider = temp.path().join("must-not-be-invoked"); + let document = temp.path().join("mcp-import.xlsx"); + let mut child = tokio::process::Command::new(binary()) + .args(["mcp", "serve", "office-native"]) + .env("A3S_OFFICECLI_EXECUTABLE", &provider) + .stdin(Stdio::piped()) + .stdout(Stdio::piped()) + .stderr(Stdio::piped()) + .kill_on_drop(true) + .spawn() + .unwrap(); + let mut stdin = child.stdin.take().unwrap(); + let mut stdout = BufReader::new(child.stdout.take().unwrap()); + let mut stderr = child.stderr.take().unwrap(); + + request( + &mut stdin, + &mut stdout, + serde_json::json!({ + "jsonrpc": "2.0", + "id": 1, + "method": "initialize", + "params": { + "protocolVersion": "2025-06-18", + "capabilities": {}, + "clientInfo": { "name": "office-import-test", "version": "1" } + } + }), + TIMEOUT, + ) + .await; + stdin + .write_all( + b"{\"jsonrpc\":\"2.0\",\"method\":\"notifications/initialized\",\"params\":{}}\n", + ) + .await + .unwrap(); + stdin.flush().await.unwrap(); + + let tools = request( + &mut stdin, + &mut stdout, + serde_json::json!({"jsonrpc":"2.0","id":2,"method":"tools/list","params":{}}), + TIMEOUT, + ) + .await; + let schema = tools["result"]["tools"] + .as_array() + .unwrap() + .iter() + .find(|tool| tool["name"] == "office_apply_batch") + .unwrap()["inputSchema"] + .to_string(); + assert!(schema.contains("import-spreadsheet-delimited")); + assert!(schema.contains("set-spreadsheet-frozen-pane")); + assert!(schema.contains("startCell")); + + let created = call( + &mut stdin, + &mut stdout, + 3, + "office_create", + serde_json::json!({"session":"workbook","file":document}), + TIMEOUT, + ) + .await; + assert_ne!(created["result"]["isError"], true, "{created}"); + + let applied = call( + &mut stdin, + &mut stdout, + 4, + "office_apply_batch", + serde_json::json!({ + "session": "workbook", + "mutations": [{ + "operation": "import-spreadsheet-delimited", + "sheet": "/Sheet1", + "import": { + "content": "Name,Amount,Date\nAlpha,42,2026-07-17\nBeta,TRUE,2026-07-18", + "format": "csv", + "header": true, + "startCell": "B2" + } + }] + }), + TIMEOUT, + ) + .await; + assert_ne!(applied["result"]["isError"], true, "{applied}"); + let content = &applied["result"]["structuredContent"]; + assert_eq!(content["persisted"], false); + assert_eq!(content["result"]["paths"][0], "/Sheet1/B2:D4"); + assert_eq!(content["result"]["spreadsheetImports"][0]["rowCount"], 3); + assert_eq!( + content["result"]["spreadsheetImports"][0]["freezePath"], + "/Sheet1/freeze" + ); + + let unsaved_date = call( + &mut stdin, + &mut stdout, + 5, + "office_get", + serde_json::json!({"session":"workbook","path":"/Sheet1/D3","depth":0}), + TIMEOUT, + ) + .await; + assert_eq!( + unsaved_date["result"]["structuredContent"]["node"]["format"]["numberFormat"], + "yyyy-mm-dd" + ); + let unsaved_freeze = call( + &mut stdin, + &mut stdout, + 6, + "office_get", + serde_json::json!({"session":"workbook","path":"/Sheet1/freeze","depth":0}), + TIMEOUT, + ) + .await; + assert_eq!( + unsaved_freeze["result"]["structuredContent"]["node"]["format"]["topLeftCell"], + "B3" + ); + let unsaved_filter = call( + &mut stdin, + &mut stdout, + 7, + "office_get", + serde_json::json!({"session":"workbook","path":"/Sheet1/autofilter","depth":0}), + TIMEOUT, + ) + .await; + assert_eq!( + unsaved_filter["result"]["structuredContent"]["node"]["format"]["ref"], + "B2:D4" + ); + let pane_set = call( + &mut stdin, + &mut stdout, + 70, + "office_apply_batch", + serde_json::json!({ + "session": "workbook", + "mutations": [{ + "operation": "set-spreadsheet-frozen-pane", + "sheet": "/Sheet1", + "pane": { + "frozenRows": 1, + "frozenColumns": 1, + "topLeftCell": "B2" + } + }] + }), + TIMEOUT, + ) + .await; + assert_ne!(pane_set["result"]["isError"], true, "{pane_set}"); + let replaced_pane = call( + &mut stdin, + &mut stdout, + 71, + "office_get", + serde_json::json!({"session":"workbook","path":"/Sheet1/freeze","depth":0}), + TIMEOUT, + ) + .await; + assert_eq!( + replaced_pane["result"]["structuredContent"]["node"]["format"]["frozenRows"], + "1" + ); + assert_eq!( + replaced_pane["result"]["structuredContent"]["node"]["format"]["frozenColumns"], + "1" + ); + assert_eq!( + replaced_pane["result"]["structuredContent"]["node"]["format"]["topLeftCell"], + "B2" + ); + + call( + &mut stdin, + &mut stdout, + 8, + "office_save", + serde_json::json!({"session":"workbook"}), + TIMEOUT, + ) + .await; + call( + &mut stdin, + &mut stdout, + 9, + "office_close", + serde_json::json!({"session":"workbook"}), + TIMEOUT, + ) + .await; + let reopened = call( + &mut stdin, + &mut stdout, + 10, + "office_open", + serde_json::json!({"session":"reopened","file":document}), + TIMEOUT, + ) + .await; + assert_ne!(reopened["result"]["isError"], true, "{reopened}"); + let persisted = call( + &mut stdin, + &mut stdout, + 11, + "office_get", + serde_json::json!({"session":"reopened","path":"/Sheet1/B4","depth":0}), + TIMEOUT, + ) + .await; + assert_eq!( + persisted["result"]["structuredContent"]["node"]["text"], + "Beta" + ); + let persisted_pane = call( + &mut stdin, + &mut stdout, + 72, + "office_get", + serde_json::json!({"session":"reopened","path":"/Sheet1/freeze","depth":0}), + TIMEOUT, + ) + .await; + assert_eq!( + persisted_pane["result"]["structuredContent"]["node"]["format"]["topLeftCell"], + "B2" + ); + let removed = call( + &mut stdin, + &mut stdout, + 12, + "office_apply_batch", + serde_json::json!({ + "session":"reopened", + "mutations":[{"operation":"remove","path":"/Sheet1/freeze"}] + }), + TIMEOUT, + ) + .await; + assert_ne!(removed["result"]["isError"], true, "{removed}"); + call( + &mut stdin, + &mut stdout, + 13, + "office_close", + serde_json::json!({"session":"reopened","discard":true}), + TIMEOUT, + ) + .await; + + drop(stdin); + let status = tokio::time::timeout(TIMEOUT, child.wait()) + .await + .unwrap() + .unwrap(); + assert!(status.success()); + let mut diagnostics = Vec::new(); + stderr.read_to_end(&mut diagnostics).await.unwrap(); + assert!( + diagnostics.is_empty(), + "{}", + String::from_utf8_lossy(&diagnostics) + ); + assert!(document.is_file()); + assert!(!provider.exists()); +} + +async fn call( + stdin: &mut tokio::process::ChildStdin, + stdout: &mut BufReader, + id: u32, + name: &str, + arguments: serde_json::Value, + timeout: Duration, +) -> serde_json::Value { + request( + stdin, + stdout, + serde_json::json!({ + "jsonrpc": "2.0", + "id": id, + "method": "tools/call", + "params": { "name": name, "arguments": arguments } + }), + timeout, + ) + .await +} + +async fn request( + stdin: &mut tokio::process::ChildStdin, + stdout: &mut BufReader, + value: serde_json::Value, + timeout: Duration, +) -> serde_json::Value { + let mut encoded = serde_json::to_vec(&value).unwrap(); + encoded.push(b'\n'); + stdin.write_all(&encoded).await.unwrap(); + stdin.flush().await.unwrap(); + let mut line = String::new(); + tokio::time::timeout(timeout, stdout.read_line(&mut line)) + .await + .unwrap() + .unwrap(); + assert!(!line.is_empty()); + serde_json::from_str(&line).unwrap() +} From 89513fc05261b229899c498a6bfd83adfa60d188 Mon Sep 17 00:00:00 2001 From: RoyLin Date: Sun, 19 Jul 2026 11:23:54 +0800 Subject: [PATCH 6/9] feat(ocr): install PP-OCRv6 models on first extraction --- README.md | 34 +++-- crates/ocr/README.md | 11 +- crates/ocr/skills/a3s-use-ocr/SKILL.md | 27 ++-- crates/ocr/src/assets.rs | 12 +- crates/ocr/src/cli.rs | 2 +- crates/ocr/src/client.rs | 17 ++- crates/ocr/src/install.rs | 190 +++++++++++++++++++++++++ crates/ocr/src/lib.rs | 4 +- crates/ocr/src/mcp.rs | 43 +++++- docs/architecture.md | 12 +- tests/cli.rs | 52 +++++++ 11 files changed, 363 insertions(+), 41 deletions(-) diff --git a/README.md b/README.md index c55a4662..8f6ee587 100644 --- a/README.md +++ b/README.md @@ -104,6 +104,7 @@ a3s use mcp serve office # Built-in local PP-OCRv6. a3s use ocr doctor --json +# The first extraction installs the pinned models when needed and allowed. a3s use ocr extract ./scan.png --json a3s use mcp serve ocr ``` @@ -153,7 +154,8 @@ Every domain argument accepted by `a3s use ...` can also be passed directly to SHA-256 for every `SKILL.md`, allowing consumers to verify the exact bytes before loading them - **Managed Provider Safety**: Require explicit installation authority, bounded - downloads, approved HTTPS origins, receipts, staging, and atomic activation + first-use policy, bounded downloads, approved HTTPS origins, receipts, + staging, and atomic activation - **Structured Automation**: Return versioned `--json` documents and typed error codes while retaining native process status and streams for delegated commands - **Component Ownership**: Remove only A3S-managed provider or package files; @@ -167,7 +169,7 @@ Every domain argument accepted by `a3s use ...` can also be passed directly to | Browser | Built in | Full Browser vocabulary | A3S Use standard MCP server | Six packaged Browser Skills | A3S Use | | Office | Built in | Stable Office vocabulary | Typed native preview plus OfficeCLI compatibility server | Packaged `a3s-use-office` Skill | A3S Use native engine; OfficeCLI compatibility in 0.1.x | | Box | Reserved built-in route | Native A3S Box vocabulary | — | — | Umbrella A3S CLI | -| OCR | Built in | Doctor and typed image extraction | `ocr_doctor` and `ocr_extract` | One local PP-OCRv6 Skill | A3S Use process with ONNX Runtime | +| OCR | Built in | Doctor and first-use typed image extraction | `ocr_doctor`, confirmed `ocr_install`, and `ocr_extract` | One local PP-OCRv6 Skill | A3S Use process with ONNX Runtime | | Science | External `a3s/science` package | Source-specific retrieval commands | 13 typed `science_*` tools | One research workflow Skill | Science extension process | | External domain | Installed extension | Optional native executable | Optional standard MCP server | Optional `SKILL.md` | Extension package plus A3S Use lifecycle | @@ -1743,17 +1745,18 @@ compatibility scope, safety invariants, delivery gates, and migration plan. ## OCR `a3s-use-ocr` implements the reserved built-in `ocr` route. The default Use -release packages its `a3s-use-ocr` Skill and exposes `ocr_doctor` plus -`ocr_extract` over standard MCP, so a resident A3S Code session receives -`mcp__use_ocr__*` without installing a separate extension. +release packages its `a3s-use-ocr` Skill and exposes `ocr_doctor`, +`ocr_install`, and `ocr_extract` over standard MCP, so a resident A3S Code +session receives `mcp__use_ocr__*` without installing a separate extension. OCR has one backend: the pinned `PP-OCRv6_small` detection and recognition -models running locally through ONNX Runtime. Release archives package those -models; `a3s install use/ocr` explicitly installs or repairs the same pinned -bundle when needed. Supported inputs are bounded local PNG, JPEG, WebP, GIF, -BMP, and TIFF files. The result binds the canonical source path, media type, -byte length, and SHA-256 alongside text, recognition/detection confidence, -polygons, and bounding boxes. +models running locally through ONNX Runtime. The first CLI extraction installs +or repairs the fixed-size, SHA-256-pinned official model archives when +networking and first-use installation are allowed. `a3s install use/ocr` +prepares the same bundle explicitly. Supported inputs are bounded local PNG, +JPEG, WebP, GIF, BMP, and TIFF files. The result binds the canonical source +path, media type, byte length, and SHA-256 alongside text, +recognition/detection confidence, polygons, and bounding boxes. The pipeline decodes and normalizes the image, runs `PP-OCRv6_small_det`, applies DB post-processing and reading-order sorting, @@ -1767,9 +1770,12 @@ a3s use ocr extract ./scan.png --json a3s use mcp serve ocr ``` -A3S Code may first-use install the verified parent Use release. A missing or -damaged managed model bundle is repaired explicitly with -`a3s install use/ocr`; the Code `use` worker never installs it implicitly. +A3S Code may first-use install the verified parent Use release. When OCR models +are missing or damaged, the Code `use` worker requests the bounded +`ocr_install` MCP tool. The parent TUI must confirm that network mutation before +the worker continues to extraction. `--offline`, `A3S_OFFLINE=1`, and +`A3S_NO_AUTO_INSTALL=1` prohibit first-use model installation. Diagnostics stay +read-only, and `a3s install use/ocr` remains available for explicit preparation. See the [OCR crate](crates/ocr/README.md) for model resolution, the inference workflow, and input boundaries. diff --git a/crates/ocr/README.md b/crates/ocr/README.md index a3b6e624..4a7f5251 100644 --- a/crates/ocr/README.md +++ b/crates/ocr/README.md @@ -11,8 +11,9 @@ There is one OCR provider: - engine: `onnx-runtime` - model bundle: `PP-OCRv6_small` -The release packages the pinned detection and recognition models. If the model -bundle is absent or damaged, install or repair it explicitly: +The first extraction installs or repairs the pinned detection and recognition +models when networking and first-use installation are allowed. Prepare them +explicitly when deterministic startup or offline work is required: ```bash a3s install use/ocr @@ -47,5 +48,11 @@ a3s use ocr extract ./scan.png --json a3s use mcp serve ocr ``` +`doctor` is read-only and never downloads anything. Direct CLI extraction +prepares missing or damaged A3S-managed models automatically. Through MCP, the +`use` worker calls the separate `ocr_install` mutation, which must pass parent +confirmation before extraction continues. `A3S_OFFLINE=1` and +`A3S_NO_AUTO_INSTALL=1` prohibit this first-use download. + Supported inputs are bounded local PNG, JPEG, WebP, GIF, BMP, and TIFF files. URLs and PDF rasterization are outside this crate. diff --git a/crates/ocr/skills/a3s-use-ocr/SKILL.md b/crates/ocr/skills/a3s-use-ocr/SKILL.md index 0a876eaf..bbab0597 100644 --- a/crates/ocr/skills/a3s-use-ocr/SKILL.md +++ b/crates/ocr/skills/a3s-use-ocr/SKILL.md @@ -6,26 +6,28 @@ description: Extract text and layout evidence from local image files through the # A3S Use OCR Use the host-provided A3S Use surface. In an A3S Code `use` worker, call -`mcp__use_ocr__ocr_doctor` and `mcp__use_ocr__ocr_extract` directly. The host -owns the MCP process; do not run a shell command, install models, or read the -file through another tool. +`mcp__use_ocr__ocr_doctor`, `mcp__use_ocr__ocr_install`, and +`mcp__use_ocr__ocr_extract` directly. The host owns the MCP process; do not run +a shell command or read the file through another tool. ## Workflow 1. Call `mcp__use_ocr__ocr_doctor`. -2. Confirm that `pp-ocr-v6`, `onnx-runtime`, and `PP-OCRv6_small` are ready. -3. Call `mcp__use_ocr__ocr_extract` with the exact local image path from the +2. If the pinned model bundle is missing or broken, call + `mcp__use_ocr__ocr_install`. This bounded network mutation must pass the + parent TUI confirmation. Do not replace it with a shell installation. +3. Confirm that `pp-ocr-v6`, `onnx-runtime`, and `PP-OCRv6_small` are ready. +4. Call `mcp__use_ocr__ocr_extract` with the exact local image path from the task. -4. Preserve the returned source path, media type, size, and SHA-256. Treat the +5. Preserve the returned source path, media type, size, and SHA-256. Treat the decoded text, recognition/detection confidence, polygons, and bounding boxes as OCR evidence rather than verified source text. The engine runs detection, reading-order sorting, perspective crop correction, tall-crop rotation, recognition, and CTC decoding locally. It does not require -Python or PaddlePaddle and never sends the source image off the device. If the -doctor reports missing or damaged models, return its typed error and explicit -`a3s install use/ocr` suggestion to the parent; never install or repair models -from inside the `use` worker. +Python or PaddlePaddle and never sends the source image off the device. Offline +mode and `A3S_NO_AUTO_INSTALL=1` prohibit the bounded installer; return that +typed policy failure to the parent instead of attempting a fallback. In a CLI-only host, equivalent commands are: @@ -34,6 +36,11 @@ a3s use ocr doctor --json a3s use ocr extract "$IMAGE" --json ``` +The first extract automatically installs or repairs the pinned models when +networking and first-use installation are allowed. `doctor` remains read-only +and never downloads anything. `a3s install use/ocr` is available for explicit +preparation. + `a3s-use-ocr` accepts the same arguments when invoked as a standalone development binary. diff --git a/crates/ocr/src/assets.rs b/crates/ocr/src/assets.rs index e60e23fa..e14ea852 100644 --- a/crates/ocr/src/assets.rs +++ b/crates/ocr/src/assets.rs @@ -7,7 +7,7 @@ use crate::config::{load_detection, load_recognition, MODEL_FAMILY}; pub(crate) const RECEIPT_FILE: &str = ".a3s-ppocr-v6.json"; -#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)] +#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize, schemars::JsonSchema)] #[serde(rename_all = "kebab-case")] pub enum OcrInstallSource { Environment, @@ -16,7 +16,7 @@ pub enum OcrInstallSource { Missing, } -#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize, schemars::JsonSchema)] #[serde(rename_all = "camelCase")] pub struct OcrRuntimeStatus { pub available: bool, @@ -95,7 +95,9 @@ pub(crate) fn resolve_model_assets() -> UseResult { "use.ocr.model_missing", format!("The local {MODEL_FAMILY} model bundle is not installed."), ) - .with_suggestion("Run 'a3s install use/ocr'.") + .with_suggestion( + "Call the bounded ocr_install MCP tool, or run 'a3s install use/ocr' explicitly.", + ) .with_detail("source", "missing") .with_detail("modelDir", managed.display().to_string())) } @@ -222,7 +224,9 @@ fn path_exists(path: &Path) -> UseResult { fn model_error(source: OcrInstallSource, root: &Path, message: impl Into) -> UseError { UseError::new("use.ocr.model_invalid", message) - .with_suggestion("Run 'a3s install use/ocr --force' to restore the pinned PP-OCRv6 bundle.") + .with_suggestion( + "Call the bounded ocr_install MCP tool, or run 'a3s install use/ocr --force' explicitly.", + ) .with_detail("source", source_name(source)) .with_detail("modelDir", root.display().to_string()) } diff --git a/crates/ocr/src/cli.rs b/crates/ocr/src/cli.rs index f5e3e690..e5550740 100644 --- a/crates/ocr/src/cli.rs +++ b/crates/ocr/src/cli.rs @@ -118,7 +118,7 @@ pub async fn run(args: Vec) -> UseResult { match cli.command { Command::Doctor => CommandOutput::data(client.diagnostic()), Command::Extract { path } => { - CommandOutput::data(client.extract(OcrRequest { path }).await?) + CommandOutput::data(client.extract_with_first_use(OcrRequest { path }).await?) } Command::Serve { .. } => Err(UseError::new( "use.ocr.command_invalid", diff --git a/crates/ocr/src/client.rs b/crates/ocr/src/client.rs index 9dc3e33c..7cf6f10a 100644 --- a/crates/ocr/src/client.rs +++ b/crates/ocr/src/client.rs @@ -8,6 +8,7 @@ use tokio::io::AsyncReadExt; use crate::assets::{ocr_status, resolve_model_assets, OcrInstallSource}; use crate::config::MODEL_FAMILY; use crate::engine::{EngineBlock, PpOcrV6Engine}; +use crate::install::ensure_ppocr_v6_ready; use crate::models::{ OcrBlock, OcrBoundingBox, OcrDiagnostic, OcrPoint, OcrProviderKind, OcrRequest, OcrResult, }; @@ -41,7 +42,7 @@ impl OcrClient { ( Readiness::Missing, vec![ - "Run 'a3s install use/ocr' to install the pinned local model bundle." + "Call the bounded ocr_install MCP tool, or run 'a3s install use/ocr' explicitly." .to_string(), ], ) @@ -49,7 +50,7 @@ impl OcrClient { ( Readiness::Broken, vec![ - "Run 'a3s install use/ocr --force' to restore the pinned local model bundle." + "Call the bounded ocr_install MCP tool, or run 'a3s install use/ocr --force' explicitly." .to_string(), ], ) @@ -72,6 +73,18 @@ impl OcrClient { pub async fn extract(&self, request: OcrRequest) -> UseResult { let source = read_source(&request.path).await?; + self.extract_source(source).await + } + + /// Validate the local source, prepare pinned models under first-use policy, + /// and then perform the same local extraction as [`Self::extract`]. + pub async fn extract_with_first_use(&self, request: OcrRequest) -> UseResult { + let source = read_source(&request.path).await?; + ensure_ppocr_v6_ready().await?; + self.extract_source(source).await + } + + async fn extract_source(&self, source: SourceImage) -> UseResult { let loaded = Arc::clone(&self.loaded); tokio::task::spawn_blocking(move || { let image = decode_image(&source.bytes)?; diff --git a/crates/ocr/src/install.rs b/crates/ocr/src/install.rs index 91031170..75a198dc 100644 --- a/crates/ocr/src/install.rs +++ b/crates/ocr/src/install.rs @@ -66,6 +66,44 @@ struct Downloaded { sha256: String, } +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +struct AutoInstallPolicy { + offline: bool, + disabled: bool, +} + +impl AutoInstallPolicy { + fn from_env() -> UseResult { + Ok(Self { + offline: parse_environment_flag("A3S_OFFLINE", std::env::var_os("A3S_OFFLINE"))?, + disabled: parse_environment_flag( + "A3S_NO_AUTO_INSTALL", + std::env::var_os("A3S_NO_AUTO_INSTALL"), + )?, + }) + } +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +enum AutoInstallAction { + Ready, + Install, +} + +/// Ensure the pinned PP-OCRv6 bundle is ready for an actual OCR operation. +/// +/// Read-only diagnostics deliberately do not call this function. Direct OCR +/// extraction and the bounded MCP install tool use it so first use installs or +/// repairs A3S-managed models while preserving offline, no-auto-install, and +/// explicit-model-directory boundaries. +pub async fn ensure_ppocr_v6_ready() -> UseResult { + let status = ocr_status(); + match automatic_install_action(&status, AutoInstallPolicy::from_env()?)? { + AutoInstallAction::Ready => Ok(status), + AutoInstallAction::Install => install_ppocr_v6(false).await, + } +} + pub async fn install_ppocr_v6(force: bool) -> UseResult { let current = ocr_status(); if !force && current.available { @@ -659,6 +697,66 @@ fn owned_install(path: &Path) -> bool { }) } +fn automatic_install_action( + status: &OcrRuntimeStatus, + policy: AutoInstallPolicy, +) -> UseResult { + if status.available { + return Ok(AutoInstallAction::Ready); + } + if status.source == OcrInstallSource::Environment { + return Err(ocr_error( + "use.ocr.model_unreadable", + format!( + "The explicit A3S_OCR_MODEL_DIR is not usable: {}", + status.detail + ), + ) + .with_suggestion("Fix or unset A3S_OCR_MODEL_DIR before retrying OCR.")); + } + if policy.offline || policy.disabled { + let reason = if policy.offline { + "offline mode" + } else { + "A3S_NO_AUTO_INSTALL" + }; + return Err(ocr_error( + "use.ocr.auto_install_disabled", + format!( + "The local {MODEL_FAMILY} bundle is not ready and automatic installation is disabled by {reason}." + ), + ) + .with_suggestion( + "Enable first-use installation or run 'a3s install use/ocr' explicitly while online.", + ) + .with_detail("reason", reason)); + } + Ok(AutoInstallAction::Install) +} + +fn parse_environment_flag(name: &str, value: Option) -> UseResult { + let Some(value) = value else { + return Ok(false); + }; + if value.is_empty() { + return Ok(true); + } + let value = value.into_string().map_err(|_| { + ocr_error( + "use.ocr.policy_invalid", + format!("{name} must contain a valid UTF-8 boolean value."), + ) + })?; + match value.trim().to_ascii_lowercase().as_str() { + "1" | "true" | "yes" | "on" => Ok(true), + "0" | "false" | "no" | "off" => Ok(false), + _ => Err(ocr_error( + "use.ocr.policy_invalid", + format!("{name} must be a boolean value."), + )), + } +} + fn archive_error(error: impl std::fmt::Display) -> UseError { ocr_error( "use.ocr.archive_invalid", @@ -669,3 +767,95 @@ fn archive_error(error: impl std::fmt::Display) -> UseError { fn ocr_error(code: &str, message: impl Into) -> UseError { UseError::new(code, message) } + +#[cfg(test)] +mod automatic_install_tests { + use super::*; + + fn status(available: bool, source: OcrInstallSource) -> OcrRuntimeStatus { + OcrRuntimeStatus { + available, + source, + model: MODEL_FAMILY.to_string(), + model_dir: None, + managed_root: None, + detail: if available { + "ready".to_string() + } else { + "missing".to_string() + }, + } + } + + #[test] + fn ready_models_never_require_an_install() { + let action = automatic_install_action( + &status(true, OcrInstallSource::Managed), + AutoInstallPolicy { + offline: true, + disabled: true, + }, + ) + .unwrap(); + + assert_eq!(action, AutoInstallAction::Ready); + } + + #[test] + fn missing_models_install_when_first_use_mutation_is_allowed() { + let action = automatic_install_action( + &status(false, OcrInstallSource::Missing), + AutoInstallPolicy { + offline: false, + disabled: false, + }, + ) + .unwrap(); + + assert_eq!(action, AutoInstallAction::Install); + } + + #[test] + fn offline_and_no_auto_install_are_strict_boundaries() { + for policy in [ + AutoInstallPolicy { + offline: true, + disabled: false, + }, + AutoInstallPolicy { + offline: false, + disabled: true, + }, + ] { + let error = automatic_install_action(&status(false, OcrInstallSource::Missing), policy) + .unwrap_err(); + assert_eq!(error.code, "use.ocr.auto_install_disabled"); + } + } + + #[test] + fn an_invalid_explicit_model_directory_is_never_replaced_implicitly() { + let error = automatic_install_action( + &status(false, OcrInstallSource::Environment), + AutoInstallPolicy { + offline: false, + disabled: false, + }, + ) + .unwrap_err(); + + assert_eq!(error.code, "use.ocr.model_unreadable"); + } + + #[test] + fn environment_flags_follow_a3s_boolean_conventions() { + for value in [None, Some("0"), Some("false"), Some("no"), Some("off")] { + assert!(!parse_environment_flag("A3S_OFFLINE", value.map(Into::into)).unwrap()); + } + for value in [Some(""), Some("1"), Some("true"), Some("yes"), Some("on")] { + assert!(parse_environment_flag("A3S_OFFLINE", value.map(Into::into)).unwrap()); + } + let error = parse_environment_flag("A3S_OFFLINE", Some("sometimes".into())).unwrap_err(); + assert_eq!(error.code, "use.ocr.policy_invalid"); + } +} diff --git a/crates/ocr/src/lib.rs b/crates/ocr/src/lib.rs index 57486889..cc36b31f 100644 --- a/crates/ocr/src/lib.rs +++ b/crates/ocr/src/lib.rs @@ -18,7 +18,9 @@ mod preprocess; pub use assets::{ocr_status, OcrInstallSource, OcrRuntimeStatus}; pub use client::OcrClient; -pub use install::{install_ppocr_v6, repair_ppocr_v6, uninstall_managed_ppocr_v6}; +pub use install::{ + ensure_ppocr_v6_ready, install_ppocr_v6, repair_ppocr_v6, uninstall_managed_ppocr_v6, +}; pub use mcp::OcrMcpServer; pub use models::{ OcrBlock, OcrBoundingBox, OcrDiagnostic, OcrPoint, OcrProviderKind, OcrRequest, OcrResult, diff --git a/crates/ocr/src/mcp.rs b/crates/ocr/src/mcp.rs index 1a3db29f..229bb609 100644 --- a/crates/ocr/src/mcp.rs +++ b/crates/ocr/src/mcp.rs @@ -5,7 +5,10 @@ use rmcp::model::{CallToolResult, Implementation, ServerCapabilities, ServerInfo use rmcp::{tool, tool_handler, tool_router, ServerHandler, ServiceExt}; use serde::Serialize; -use crate::{OcrClient, OcrDiagnostic, OcrRequest, OcrResult, UseError, UseResult}; +use crate::{ + ensure_ppocr_v6_ready, OcrClient, OcrDiagnostic, OcrRequest, OcrResult, OcrRuntimeStatus, + UseError, UseResult, +}; #[derive(Clone)] pub struct OcrMcpServer { @@ -56,6 +59,21 @@ impl OcrMcpServer { Ok(tool_result(Ok(self.client.diagnostic()))) } + #[tool( + name = "ocr_install", + description = "Install or repair the pinned local PP-OCRv6 model bundle from its official HTTPS source with fixed size and SHA-256 checks", + output_schema = rmcp::handler::server::tool::cached_schema_for_type::(), + annotations( + read_only_hint = false, + destructive_hint = false, + idempotent_hint = true, + open_world_hint = true + ) + )] + async fn ocr_install(&self) -> Result { + Ok(tool_result(ensure_ppocr_v6_ready().await)) + } + #[tool( name = "ocr_extract", description = "Extract text, polygons, bounding boxes, and confidence from one bounded local image with PP-OCRv6; source bytes remain on this device", @@ -88,7 +106,7 @@ impl ServerHandler for OcrMcpServer { website_url: Some("https://github.com/A3S-Lab/Use".to_string()), }, instructions: Some( - "Call ocr_doctor first. Use ocr_extract only for a local image path supplied by the task. PP-OCRv6 detection and recognition run locally through ONNX Runtime and never send source bytes off device. Preserve the source SHA-256 and distinguish OCR text from verified source text." + "Call ocr_doctor first. When the pinned models are missing or broken, request ocr_install through the host confirmation path, then call ocr_extract with the local image path supplied by the task. PP-OCRv6 detection and recognition run locally through ONNX Runtime and never send source bytes off device. Preserve the source SHA-256 and distinguish OCR text from verified source text." .to_string(), ), ..Default::default() @@ -143,15 +161,20 @@ mod tests { .iter() .map(|tool| tool.name.as_ref()) .collect::>(), - ["ocr_doctor", "ocr_extract"] + ["ocr_doctor", "ocr_extract", "ocr_install"] ); let doctor = tools.iter().find(|tool| tool.name == "ocr_doctor").unwrap(); let extract = tools .iter() .find(|tool| tool.name == "ocr_extract") .unwrap(); + let install = tools + .iter() + .find(|tool| tool.name == "ocr_install") + .unwrap(); assert!(doctor.output_schema.is_some()); assert!(extract.output_schema.is_some()); + assert!(install.output_schema.is_some()); assert_eq!( doctor .annotations @@ -166,5 +189,19 @@ mod tests { .and_then(|annotations| annotations.open_world_hint), Some(false) ); + assert_eq!( + install + .annotations + .as_ref() + .and_then(|annotations| annotations.read_only_hint), + Some(false) + ); + assert_eq!( + install + .annotations + .as_ref() + .and_then(|annotations| annotations.open_world_hint), + Some(true) + ); } } diff --git a/docs/architecture.md b/docs/architecture.md index b1ffebee..4dc16a81 100644 --- a/docs/architecture.md +++ b/docs/architecture.md @@ -43,8 +43,11 @@ CLI plus standard stdio MCP without a separate extension install. The process accepts bounded local image files and binds every result to the canonical source digest. It runs the pinned `PP-OCRv6_small` detection and recognition models locally through ONNX Runtime, without Python, PaddlePaddle, a remote OCR -endpoint, or an alternate backend. Model installation and repair are explicit -`use/ocr` component operations. Both MCP tools are closed-world and read-only. +endpoint, or an alternate backend. The first CLI extraction installs or repairs +the pinned model bundle when first-use policy permits it. Standard MCP keeps +`ocr_doctor` and `ocr_extract` closed-world and read-only, while the separate +idempotent `ocr_install` network mutation requires parent confirmation. +Explicit `use/ocr` component operations remain available for preparation. `a3s-use-science` is the reference multi-surface extension. It remains a separate process and package even though its source is developed in this @@ -105,8 +108,9 @@ compatibility provider. For resident hosts, `use/office` targets the built-in `office-native` MCP server and is ready independently of OfficeCLI. A discovered OfficeCLI provider is projected separately as `use/office-compat`, targeting the standard compatibility server without carrying the native Skill. The -`use/ocr` route targets `ocr-native`; model readiness remains explicit and -never triggers a silent install. +`use/ocr` route targets `ocr-native`; model readiness remains visible. +Read-only discovery never installs models, while direct extraction uses the +first-use policy and the MCP worker uses a separately confirmed install tool. The projection contains content-bound Skill references and an MCP launch target, never executable extension code or a generic action payload. Consumers still diff --git a/tests/cli.rs b/tests/cli.rs index 68f813d2..4c1034f3 100644 --- a/tests/cli.rs +++ b/tests/cli.rs @@ -2196,6 +2196,58 @@ fn built_in_ocr_projects_the_canonical_code_route_and_skill() { assert_eq!(digest.len(), 64); } +#[cfg(feature = "ocr")] +#[test] +fn ocr_extract_honors_the_no_auto_install_boundary_for_a_valid_image() { + let temp = tempfile::tempdir().unwrap(); + let image_path = temp.path().join("input.bmp"); + std::fs::write( + &image_path, + [ + 0x42, 0x4d, 0x3a, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x36, 0x00, 0x00, 0x00, + 0x28, 0x00, 0x00, 0x00, 0x01, 0x00, 0x00, 0x00, 0x01, 0x00, 0x00, 0x00, 0x01, 0x00, + 0x18, 0x00, 0x00, 0x00, 0x00, 0x00, 0x04, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, + 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0xff, 0xff, + 0xff, 0x00, + ], + ) + .unwrap(); + let output = Command::new(binary()) + .args(["ocr", "extract", image_path.to_str().unwrap(), "--json"]) + .env("A3S_USE_OCR_HOME", temp.path().join("ocr")) + .env("A3S_NO_AUTO_INSTALL", "1") + .output() + .unwrap(); + + assert_eq!(output.status.code(), Some(1), "{output:?}"); + let value: serde_json::Value = serde_json::from_slice(&output.stdout).unwrap(); + assert_eq!(value["ok"], false); + assert_eq!(value["error"]["code"], "use.ocr.auto_install_disabled"); + assert_eq!(value["error"]["details"]["reason"], "A3S_NO_AUTO_INSTALL"); +} + +#[cfg(feature = "ocr")] +#[test] +fn ocr_extract_validates_the_source_before_first_use_installation() { + let temp = tempfile::tempdir().unwrap(); + let output = Command::new(binary()) + .args([ + "ocr", + "extract", + temp.path().join("missing.png").to_str().unwrap(), + "--json", + ]) + .env("A3S_USE_OCR_HOME", temp.path().join("ocr")) + .env("A3S_NO_AUTO_INSTALL", "1") + .output() + .unwrap(); + + assert_eq!(output.status.code(), Some(1), "{output:?}"); + let value: serde_json::Value = serde_json::from_slice(&output.stdout).unwrap(); + assert_eq!(value["ok"], false); + assert_eq!(value["error"]["code"], "use.ocr.source_unreadable"); +} + #[cfg(all(unix, feature = "extensions"))] #[test] fn explicit_extension_install_delegates_native_cli_and_preserves_status() { From 76dd9f90f2eb63a64f91c756d8ae8b7a6839f793 Mon Sep 17 00:00:00 2001 From: RoyLin Date: Sun, 19 Jul 2026 13:19:09 +0800 Subject: [PATCH 7/9] feat: prepare application runtimes on first use --- Cargo.lock | 1 + README.md | 38 +++-- crates/browser-driver/Cargo.toml | 1 + .../skills/a3s-use-browser/SKILL.md | 8 +- crates/browser-driver/src/lifecycle.rs | 159 +++++++++++++++++- crates/browser-driver/src/main.rs | 72 ++++++++ crates/core/src/lib.rs | 107 ++++++++++++ crates/ocr/src/install.rs | 92 ++-------- crates/office/skills/a3s-use-office/SKILL.md | 17 +- .../skills/a3s-use-office/references/mcp.md | 13 +- docs/architecture.md | 31 ++-- src/browser_cli.rs | 1 + src/cli.rs | 16 ++ src/first_use.rs | 145 ++++++++++++++++ src/lib.rs | 1 + src/mcp/office.rs | 21 ++- src/mcp/office/tests.rs | 15 +- tests/cli.rs | 75 ++++++++- 18 files changed, 680 insertions(+), 133 deletions(-) create mode 100644 src/first_use.rs diff --git a/Cargo.lock b/Cargo.lock index 7422d187..7ef15ef8 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -68,6 +68,7 @@ name = "a3s-use-browser-driver" version = "0.1.1" dependencies = [ "a3s-acl", + "a3s-use-core", "aes-gcm", "async-trait", "base64", diff --git a/README.md b/README.md index 8f6ee587..5b0cc735 100644 --- a/README.md +++ b/README.md @@ -166,8 +166,8 @@ Every domain argument accepted by `a3s use ...` can also be passed directly to | Domain | Origin | CLI | MCP | Skill | Runtime owner | | --- | --- | --- | --- | --- | --- | -| Browser | Built in | Full Browser vocabulary | A3S Use standard MCP server | Six packaged Browser Skills | A3S Use | -| Office | Built in | Stable Office vocabulary | Typed native preview plus OfficeCLI compatibility server | Packaged `a3s-use-office` Skill | A3S Use native engine; OfficeCLI compatibility in 0.1.x | +| Browser | Built in | Full Browser vocabulary with first-launch preparation | A3S Use standard MCP server with confirmed installer | Six packaged Browser Skills | A3S Use | +| Office | Built in | Native Office plus first-use compatibility fallback | Typed native server with confirmed compatibility installer plus OfficeCLI server | Packaged `a3s-use-office` Skill | A3S Use native engine; OfficeCLI compatibility in 0.1.x | | Box | Reserved built-in route | Native A3S Box vocabulary | — | — | Umbrella A3S CLI | | OCR | Built in | Doctor and first-use typed image extraction | `ocr_doctor`, confirmed `ocr_install`, and `ocr_extract` | One local PP-OCRv6 Skill | A3S Use process with ONNX Runtime | | Science | External `a3s/science` package | Source-specific retrieval commands | 13 typed `science_*` tools | One research workflow Skill | Science extension process | @@ -215,6 +215,7 @@ release selection and the top-level component receipt: ```bash a3s install use --source release +# Optional deterministic pre-warm; normal first use prepares these as needed. a3s install use/browser a3s install use/office a3s use doctor --json @@ -267,9 +268,9 @@ async fn main() -> Result<(), Box> { } ``` -`BrowserPoolConfig::default()` discovers an existing Chrome-compatible browser -and never authorizes a download. Select a managed provider or run an explicit -component install when A3S should own the runtime. +`BrowserPoolConfig::default()` remains non-installing for embedded callers such +as Search. Product commands validate their arguments and then prepare the same +shared managed runtime on the first local Browser launch when policy allows. ## Browser @@ -304,12 +305,16 @@ port, requires a private bearer token, has bounded idle and maximum lifetimes, and shares typed Browser session state. It is an MCP deployment, not an A3S JSON-RPC service. -Provider selection stays explicit. Discovered providers never download -software. Managed Chrome and Lightpanda installations use bounded staging and -atomic activation; Lightpanda assets require the publisher SHA-256. Chrome for -Testing does not publish an independent checksum in its current version feed, -so A3S records HTTPS provenance and the locally observed digest without claiming -publisher verification. +Provider selection stays typed. Embedded `Discovered*` providers never +download software. A direct local Browser launch is first-use authority for the +A3S product CLI; Code workers request the bounded installer through parent +confirmation. Both paths reuse a system browser or the shared A3S-managed +cache before downloading. Managed Chrome and Lightpanda installations use +bounded staging and atomic activation; Lightpanda assets require the publisher +SHA-256. Chrome for Testing does not publish an independent checksum in its +current version feed, so A3S records HTTPS provenance and the locally observed +digest without claiming publisher verification. Help, version, doctor, Skills, +profiles, and MCP server startup never install a browser. See [Agent Browser Compatibility Baseline](docs/agent-browser-parity.md) for the locked schemas, digests, runtime evidence, and promotion criteria. @@ -617,7 +622,10 @@ Other `0.1.x` commands and the default `mcp serve office` target still use a compatibility backend pinned to OfficeCLI `1.0.136`. This is a migration boundary, not a native-promotion claim. The default routes will be promoted only after mutation, fidelity, rendering, compatibility, and cross-application -interoperability gates pass. +interoperability gates pass. The first real compatibility CLI command prepares +that pinned provider when first-use policy allows. In Code, the native Office +worker requests `office_install_compat` through parent confirmation only when +the requested operation is outside the native surface. ```bash # Inspect without downloading anything. @@ -836,7 +844,8 @@ a3s use office native dump report.docx --output report.replay.json --json a3s use office native create restored.docx --json a3s use office native batch restored.docx --input report.replay.json --json -# Install the current compatibility provider explicitly. +# Optional compatibility pre-warm. The following compatibility commands also +# prepare this pinned provider on first use. a3s install use/office a3s use office get report.docx /body --json a3s use office batch report.xlsx --input updates.json --json @@ -848,7 +857,8 @@ a3s use mcp serve office-native a3s use mcp serve office ``` -The native MCP process exposes 12 typed tools: `office_validate`, +The native MCP process exposes 12 document tools plus the confirmed +`office_install_compat` compatibility installer: `office_validate`, `office_create`, `office_open`, `office_list`, `office_get`, `office_query`, `office_view`, `office_raw_xml`, `office_apply_batch`, `office_merge_template`, `office_save`, and `office_close`. It accepts no shell diff --git a/crates/browser-driver/Cargo.toml b/crates/browser-driver/Cargo.toml index c58f3d19..8f0828ca 100644 --- a/crates/browser-driver/Cargo.toml +++ b/crates/browser-driver/Cargo.toml @@ -13,6 +13,7 @@ name = "a3s-use-browser-driver" path = "src/main.rs" [dependencies] +a3s-use-core = { version = "0.1.1", path = "../core" } a3s-acl = { git = "https://github.com/A3S-Lab/ACL", rev = "6e2a6469edc0f4c61b1e588d0ace873aaf15ce22" } serde = { version = "1.0", features = ["derive"] } serde_json = "1.0" diff --git a/crates/browser-driver/skills/a3s-use-browser/SKILL.md b/crates/browser-driver/skills/a3s-use-browser/SKILL.md index ee03a73a..1e0b537e 100644 --- a/crates/browser-driver/skills/a3s-use-browser/SKILL.md +++ b/crates/browser-driver/skills/a3s-use-browser/SKILL.md @@ -16,12 +16,18 @@ Use the host surface that is already available: must obtain HITL approval before that mutation can run. - In a CLI-only agent host, use the `a3s use browser ...` commands below. -Install the built-in capability and its managed runtime when needed: +The first direct local launch installs the shared A3S-managed Chrome runtime +when no system or managed browser is available and first-use policy permits it. +Prepare it explicitly for deterministic startup or offline work: ```bash a3s install use use/browser ``` +Doctor, help, version, Skills, profiles, and MCP server startup remain +non-installing. `A3S_OFFLINE=1` and `A3S_NO_AUTO_INSTALL=1` prohibit the +first-use download. + Load the version-matched core workflow before browser automation: ```bash diff --git a/crates/browser-driver/src/lifecycle.rs b/crates/browser-driver/src/lifecycle.rs index fbc25d57..51ad3955 100644 --- a/crates/browser-driver/src/lifecycle.rs +++ b/crates/browser-driver/src/lifecycle.rs @@ -3,8 +3,16 @@ use std::path::PathBuf; use std::process::{Command, Stdio}; +use a3s_use_core::FirstUseInstallPolicy; + const USE_EXECUTABLE_ENV: &str = "A3S_USE_EXECUTABLE"; +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +enum AutoInstallAction { + Ready, + Install, +} + pub fn install(with_dependencies: bool, json: bool) -> i32 { if with_dependencies { if let Err(error) = crate::install::install_system_dependencies() { @@ -24,18 +32,89 @@ pub fn upgrade(json: bool) -> i32 { run_component_install(true, json) } +pub fn ensure_first_use_browser() -> Result<(), String> { + let available = crate::native::cdp::chrome::find_chrome().is_some(); + let explicit_invalid = explicit_browser_provider_invalid(); + let policy = FirstUseInstallPolicy::from_env() + .map_err(|error| format!("{}: {}", error.code, error.message))?; + match automatic_install_action(available, explicit_invalid, policy)? { + AutoInstallAction::Ready => Ok(()), + AutoInstallAction::Install => { + run_component_install_captured()?; + crate::native::cdp::chrome::find_chrome() + .is_some() + .then_some(()) + .ok_or_else(|| { + "use.browser.install_failed: Browser installation completed without a usable Chrome executable." + .to_string() + }) + } + } +} + +fn automatic_install_action( + available: bool, + explicit_invalid: bool, + policy: FirstUseInstallPolicy, +) -> Result { + if explicit_invalid { + return Err( + "use.browser.explicit_provider_invalid: The explicit Browser executable is not usable. Fix or unset it before retrying." + .to_string(), + ); + } + if available { + return Ok(AutoInstallAction::Ready); + } + if let Some(block) = policy.blocked_by() { + return Err(format!( + "use.browser.auto_install_disabled: No compatible browser is ready and first-use installation is disabled by {}. Run 'a3s install use/browser' explicitly while online.", + block.reason() + )); + } + Ok(AutoInstallAction::Install) +} + +fn explicit_browser_provider_invalid() -> bool { + [ + "A3S_USE_BROWSER_EXECUTABLE_PATH", + "AGENT_BROWSER_EXECUTABLE_PATH", + "A3S_BROWSER_EXECUTABLE", + "CHROME", + ] + .iter() + .filter_map(std::env::var_os) + .filter(|value| !value.is_empty()) + .map(PathBuf::from) + .any(|path| !is_usable_executable(&path)) +} + +fn is_usable_executable(path: &std::path::Path) -> bool { + let Ok(metadata) = std::fs::metadata(path) else { + return false; + }; + if !metadata.is_file() { + return false; + } + #[cfg(unix)] + { + use std::os::unix::fs::PermissionsExt; + metadata.permissions().mode() & 0o111 != 0 + } + #[cfg(not(unix))] + { + true + } +} + fn run_component_install(force: bool, json: bool) -> i32 { - let executable = match resolve_use_executable() { - Some(executable) => executable, - None => { - eprintln!( - "Cannot find a3s-use for the Browser component lifecycle. Install or repair the A3S Use package, or set {USE_EXECUTABLE_ENV}." - ); + let mut command = match component_install_command(force, json) { + Ok(command) => command, + Err(error) => { + eprintln!("{error}"); return 1; } }; - let mut command = Command::new(&executable); - command.args(component_install_args(force, json)); command .stdin(Stdio::inherit()) .stdout(Stdio::inherit()) @@ -45,13 +124,51 @@ fn run_component_install(force: bool, json: bool) -> i32 { Err(error) => { eprintln!( "Failed to launch A3S Use component lifecycle '{}': {error}", - executable.display() + command.get_program().to_string_lossy() ); 1 } } } +fn run_component_install_captured() -> Result<(), String> { + let mut command = component_install_command(false, true)?; + let output = command + .stdin(Stdio::null()) + .stdout(Stdio::piped()) + .stderr(Stdio::piped()) + .output() + .map_err(|error| { + format!( + "use.browser.install_failed: Failed to launch A3S Use component lifecycle '{}': {error}", + command.get_program().to_string_lossy() + ) + })?; + if output.status.success() { + return Ok(()); + } + let stderr = String::from_utf8_lossy(&output.stderr); + let stdout = String::from_utf8_lossy(&output.stdout); + let detail = [stderr.trim(), stdout.trim()] + .into_iter() + .find(|value| !value.is_empty()) + .unwrap_or("the component installer returned no diagnostic"); + Err(format!( + "use.browser.install_failed: Browser first-use installation failed: {detail}" + )) +} + +fn component_install_command(force: bool, json: bool) -> Result { + let executable = resolve_use_executable().ok_or_else(|| { + format!( + "Cannot find a3s-use for the Browser component lifecycle. Install or repair the A3S Use package, or set {USE_EXECUTABLE_ENV}." + ) + })?; + let mut command = Command::new(executable); + command.args(component_install_args(force, json)); + Ok(command) +} + fn component_install_args(force: bool, json: bool) -> Vec<&'static str> { let mut arguments = vec!["component", "install", "browser"]; if force { @@ -124,6 +241,30 @@ mod tests { ); } + #[test] + fn first_use_installs_only_when_the_runtime_is_missing_and_policy_allows_it() { + assert_eq!( + automatic_install_action(true, false, FirstUseInstallPolicy::new(true, true)).unwrap(), + AutoInstallAction::Ready + ); + assert_eq!( + automatic_install_action(false, false, FirstUseInstallPolicy::new(false, false)) + .unwrap(), + AutoInstallAction::Install + ); + for policy in [ + FirstUseInstallPolicy::new(true, false), + FirstUseInstallPolicy::new(false, true), + ] { + let error = automatic_install_action(false, false, policy).unwrap_err(); + assert!(error.starts_with("use.browser.auto_install_disabled:")); + } + let explicit = + automatic_install_action(false, true, FirstUseInstallPolicy::new(false, false)) + .unwrap_err(); + assert!(explicit.starts_with("use.browser.explicit_provider_invalid:")); + } + #[test] fn explicit_lifecycle_executable_must_be_a_file() { let temp = tempfile::tempdir().unwrap(); diff --git a/crates/browser-driver/src/main.rs b/crates/browser-driver/src/main.rs index db3e20d7..a0c02392 100644 --- a/crates/browser-driver/src/main.rs +++ b/crates/browser-driver/src/main.rs @@ -156,6 +156,33 @@ fn incompatible_launch_mode_error(flags: &Flags) -> Option<&'static str> { None } +fn uses_external_browser(flags: &Flags) -> bool { + flags.executable_path.is_some() + || flags.provider.is_some() + || flags.cdp.is_some() + || flags.auto_connect +} + +fn command_starts_local_browser(command: &serde_json::Value, flags: &Flags) -> bool { + if uses_external_browser(flags) + || command.get("cdpUrl").is_some() + || command.get("cdpPort").is_some() + { + return false; + } + matches!( + command.get("action").and_then(|value| value.as_str()), + Some("launch" | "navigate" | "batch" | "diff_url") + ) +} + +fn prepare_first_use_browser(flags: &Flags) -> Result<(), String> { + if uses_external_browser(flags) { + return Ok(()); + } + lifecycle::ensure_first_use_browser() +} + fn should_send_local_launch_config(flags: &Flags) -> bool { (flags.headed || flags.cli_headed @@ -1111,6 +1138,14 @@ fn main() { } else { None }; + if let Err(error) = prepare_first_use_browser(&flags) { + if flags.json { + print_json_error(error); + } else { + eprintln!("{} {}", color::error_indicator(), error); + } + exit(1); + } chat::run_chat(&flags, message); return; } @@ -1237,6 +1272,17 @@ fn main() { exit(1); } + if command_starts_local_browser(&cmd, &flags) { + if let Err(error) = prepare_first_use_browser(&flags) { + if flags.json { + print_json_error(error); + } else { + eprintln!("{} {}", color::error_indicator(), error); + } + exit(1); + } + } + // Parse proxy URL to separate server from credentials for the daemon. let (proxy_server, proxy_username, proxy_password) = if let Some(ref proxy_str) = flags.proxy { let parsed = parse_proxy(proxy_str); @@ -1959,6 +2005,32 @@ mod tests { flags } + #[test] + fn first_use_preparation_is_limited_to_local_browser_launches() { + let flags = neutral_launch_config_flags(); + for action in ["launch", "navigate", "batch", "diff_url"] { + assert!(command_starts_local_browser( + &json!({ "action": action }), + &flags + )); + } + assert!(!command_starts_local_browser( + &json!({ "action": "snapshot" }), + &flags + )); + assert!(!command_starts_local_browser( + &json!({ "action": "launch", "cdpUrl": "http://127.0.0.1:9222" }), + &flags + )); + + let mut external = neutral_launch_config_flags(); + external.executable_path = Some("/explicit/chrome".to_string()); + assert!(!command_starts_local_browser( + &json!({ "action": "navigate" }), + &external + )); + } + #[test] fn test_attach_allowed_domains_to_launch_command() { let mut flags = neutral_launch_config_flags(); diff --git a/crates/core/src/lib.rs b/crates/core/src/lib.rs index 80fb2e7e..fcc4eaac 100644 --- a/crates/core/src/lib.rs +++ b/crates/core/src/lib.rs @@ -1,4 +1,5 @@ use std::collections::BTreeMap; +use std::ffi::OsString; use std::fmt; use std::path::PathBuf; @@ -120,6 +121,86 @@ impl std::error::Error for UseError {} pub type UseResult = Result; +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum FirstUseInstallBlock { + Offline, + Disabled, +} + +impl FirstUseInstallBlock { + pub const fn reason(self) -> &'static str { + match self { + Self::Offline => "offline mode", + Self::Disabled => "A3S_NO_AUTO_INSTALL", + } + } +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub struct FirstUseInstallPolicy { + offline: bool, + disabled: bool, +} + +impl FirstUseInstallPolicy { + pub const fn new(offline: bool, disabled: bool) -> Self { + Self { offline, disabled } + } + + pub fn from_env() -> UseResult { + Self::from_values( + std::env::var_os("A3S_OFFLINE"), + std::env::var_os("A3S_NO_AUTO_INSTALL"), + ) + } + + pub const fn blocked_by(self) -> Option { + if self.offline { + Some(FirstUseInstallBlock::Offline) + } else if self.disabled { + Some(FirstUseInstallBlock::Disabled) + } else { + None + } + } + + pub const fn allows_install(self) -> bool { + self.blocked_by().is_none() + } + + fn from_values(offline: Option, disabled: Option) -> UseResult { + Ok(Self { + offline: parse_environment_boolean("A3S_OFFLINE", offline)?, + disabled: parse_environment_boolean("A3S_NO_AUTO_INSTALL", disabled)?, + }) + } +} + +fn parse_environment_boolean(name: &'static str, value: Option) -> UseResult { + let Some(value) = value else { + return Ok(false); + }; + if value.is_empty() { + return Ok(true); + } + let value = value.into_string().map_err(|_| { + UseError::new( + "use.first_use.policy_invalid", + format!("{name} must contain a valid UTF-8 boolean value."), + ) + .with_detail("variable", name) + })?; + match value.trim().to_ascii_lowercase().as_str() { + "1" | "true" | "yes" | "on" => Ok(true), + "0" | "false" | "no" | "off" => Ok(false), + _ => Err(UseError::new( + "use.first_use.policy_invalid", + format!("{name} must be a boolean value."), + ) + .with_detail("variable", name)), + } +} + #[cfg(test)] mod tests { use super::*; @@ -143,4 +224,30 @@ mod tests { .unwrap() .contains("a3s install")); } + + #[test] + fn first_use_policy_uses_a3s_boolean_conventions() { + for value in [None, Some("0"), Some("false"), Some("no"), Some("off")] { + let policy = FirstUseInstallPolicy::from_values(value.map(Into::into), None).unwrap(); + assert!(policy.allows_install()); + } + for value in [Some(""), Some("1"), Some("true"), Some("yes"), Some("on")] { + let policy = FirstUseInstallPolicy::from_values(value.map(Into::into), None).unwrap(); + assert_eq!(policy.blocked_by(), Some(FirstUseInstallBlock::Offline)); + } + } + + #[test] + fn offline_policy_takes_precedence_over_no_auto_install() { + let policy = + FirstUseInstallPolicy::from_values(Some("1".into()), Some("1".into())).unwrap(); + assert_eq!(policy.blocked_by(), Some(FirstUseInstallBlock::Offline)); + } + + #[test] + fn invalid_first_use_policy_is_typed() { + let error = FirstUseInstallPolicy::from_values(Some("sometimes".into()), None).unwrap_err(); + assert_eq!(error.code, "use.first_use.policy_invalid"); + assert_eq!(error.details["variable"], "A3S_OFFLINE"); + } } diff --git a/crates/ocr/src/install.rs b/crates/ocr/src/install.rs index 75a198dc..17ecc8b8 100644 --- a/crates/ocr/src/install.rs +++ b/crates/ocr/src/install.rs @@ -3,7 +3,7 @@ use std::io::{Read, Write}; use std::path::{Component, Path, PathBuf}; use std::sync::atomic::{AtomicU64, Ordering}; -use a3s_use_core::{UseError, UseResult}; +use a3s_use_core::{FirstUseInstallPolicy, UseError, UseResult}; use fs2::FileExt; use serde::{Deserialize, Serialize}; use sha2::{Digest, Sha256}; @@ -66,24 +66,6 @@ struct Downloaded { sha256: String, } -#[derive(Debug, Clone, Copy, PartialEq, Eq)] -struct AutoInstallPolicy { - offline: bool, - disabled: bool, -} - -impl AutoInstallPolicy { - fn from_env() -> UseResult { - Ok(Self { - offline: parse_environment_flag("A3S_OFFLINE", std::env::var_os("A3S_OFFLINE"))?, - disabled: parse_environment_flag( - "A3S_NO_AUTO_INSTALL", - std::env::var_os("A3S_NO_AUTO_INSTALL"), - )?, - }) - } -} - #[derive(Debug, Clone, Copy, PartialEq, Eq)] enum AutoInstallAction { Ready, @@ -98,7 +80,7 @@ enum AutoInstallAction { /// explicit-model-directory boundaries. pub async fn ensure_ppocr_v6_ready() -> UseResult { let status = ocr_status(); - match automatic_install_action(&status, AutoInstallPolicy::from_env()?)? { + match automatic_install_action(&status, FirstUseInstallPolicy::from_env()?)? { AutoInstallAction::Ready => Ok(status), AutoInstallAction::Install => install_ppocr_v6(false).await, } @@ -699,7 +681,7 @@ fn owned_install(path: &Path) -> bool { fn automatic_install_action( status: &OcrRuntimeStatus, - policy: AutoInstallPolicy, + policy: FirstUseInstallPolicy, ) -> UseResult { if status.available { return Ok(AutoInstallAction::Ready); @@ -714,12 +696,8 @@ fn automatic_install_action( ) .with_suggestion("Fix or unset A3S_OCR_MODEL_DIR before retrying OCR.")); } - if policy.offline || policy.disabled { - let reason = if policy.offline { - "offline mode" - } else { - "A3S_NO_AUTO_INSTALL" - }; + if let Some(block) = policy.blocked_by() { + let reason = block.reason(); return Err(ocr_error( "use.ocr.auto_install_disabled", format!( @@ -734,29 +712,6 @@ fn automatic_install_action( Ok(AutoInstallAction::Install) } -fn parse_environment_flag(name: &str, value: Option) -> UseResult { - let Some(value) = value else { - return Ok(false); - }; - if value.is_empty() { - return Ok(true); - } - let value = value.into_string().map_err(|_| { - ocr_error( - "use.ocr.policy_invalid", - format!("{name} must contain a valid UTF-8 boolean value."), - ) - })?; - match value.trim().to_ascii_lowercase().as_str() { - "1" | "true" | "yes" | "on" => Ok(true), - "0" | "false" | "no" | "off" => Ok(false), - _ => Err(ocr_error( - "use.ocr.policy_invalid", - format!("{name} must be a boolean value."), - )), - } -} - fn archive_error(error: impl std::fmt::Display) -> UseError { ocr_error( "use.ocr.archive_invalid", @@ -791,10 +746,7 @@ mod automatic_install_tests { fn ready_models_never_require_an_install() { let action = automatic_install_action( &status(true, OcrInstallSource::Managed), - AutoInstallPolicy { - offline: true, - disabled: true, - }, + FirstUseInstallPolicy::new(true, true), ) .unwrap(); @@ -805,10 +757,7 @@ mod automatic_install_tests { fn missing_models_install_when_first_use_mutation_is_allowed() { let action = automatic_install_action( &status(false, OcrInstallSource::Missing), - AutoInstallPolicy { - offline: false, - disabled: false, - }, + FirstUseInstallPolicy::new(false, false), ) .unwrap(); @@ -818,14 +767,8 @@ mod automatic_install_tests { #[test] fn offline_and_no_auto_install_are_strict_boundaries() { for policy in [ - AutoInstallPolicy { - offline: true, - disabled: false, - }, - AutoInstallPolicy { - offline: false, - disabled: true, - }, + FirstUseInstallPolicy::new(true, false), + FirstUseInstallPolicy::new(false, true), ] { let error = automatic_install_action(&status(false, OcrInstallSource::Missing), policy) .unwrap_err(); @@ -837,25 +780,10 @@ mod automatic_install_tests { fn an_invalid_explicit_model_directory_is_never_replaced_implicitly() { let error = automatic_install_action( &status(false, OcrInstallSource::Environment), - AutoInstallPolicy { - offline: false, - disabled: false, - }, + FirstUseInstallPolicy::new(false, false), ) .unwrap_err(); assert_eq!(error.code, "use.ocr.model_unreadable"); } - - #[test] - fn environment_flags_follow_a3s_boolean_conventions() { - for value in [None, Some("0"), Some("false"), Some("no"), Some("off")] { - assert!(!parse_environment_flag("A3S_OFFLINE", value.map(Into::into)).unwrap()); - } - for value in [Some(""), Some("1"), Some("true"), Some("yes"), Some("on")] { - assert!(parse_environment_flag("A3S_OFFLINE", value.map(Into::into)).unwrap()); - } - let error = parse_environment_flag("A3S_OFFLINE", Some("sometimes".into())).unwrap_err(); - assert_eq!(error.code, "use.ocr.policy_invalid"); - } } diff --git a/crates/office/skills/a3s-use-office/SKILL.md b/crates/office/skills/a3s-use-office/SKILL.md index 0cb975f0..37aa35bb 100644 --- a/crates/office/skills/a3s-use-office/SKILL.md +++ b/crates/office/skills/a3s-use-office/SKILL.md @@ -13,12 +13,13 @@ Use the host surface that is already available: - In an A3S Code `use` worker, call the available `mcp__use_office__*` tools directly. The host has already started the native - MCP server and owns its lifecycle; do not run shell commands or install a - provider. + MCP server and owns its lifecycle; do not run shell commands. - If a requested operation is absent from the native tools, use an available `mcp__use_office_compat__*` tool only as an explicit compatibility fallback. - If that route is absent, report the missing capability instead of installing, - repairing, or falling back to a shell. + If that route is absent, request + `mcp__use_office__office_install_compat`. This bounded network mutation must + pass parent TUI confirmation. Use the compatibility route only after the host + projects it; never replace the installer with a shell command. - In a CLI-only agent host, use the `a3s use office native ...` commands below. ## Workflow @@ -74,14 +75,16 @@ when an agent must bound its lifetime. - In an A3S Code `use` worker, use `mcp__use_office__*` and keep the returned Office session ID stable until the document is saved and closed. - Use `mcp__use_office_compat__*` only when the native route lacks the requested - operation and the compatibility tools are actually present. + operation. If the tools are absent, request + `mcp__use_office__office_install_compat` through parent confirmation first. - Use `a3s use office native ... --json` for local automation and scripts. - Use `a3s use mcp serve office-native` for typed, stateful agent sessions. Read [references/mcp.md](references/mcp.md) before using its session tools. - Use the typed Rust API when embedding Office behavior in Rust. - Use `a3s use office ...` only for an operation absent from the native route. - Check `a3s use office doctor --json` first. Never install or repair the - compatibility provider without explicit user authority. + Check `a3s use office doctor --json` first. The first real compatibility + command prepares the pinned OfficeCLI provider when policy allows it; help, + version, doctor, Skills, and native commands remain non-installing. `a3s-use` accepts the same arguments when the umbrella `a3s` executable is not available. diff --git a/crates/office/skills/a3s-use-office/references/mcp.md b/crates/office/skills/a3s-use-office/references/mcp.md index 7d253d56..dc7f791f 100644 --- a/crates/office/skills/a3s-use-office/references/mcp.md +++ b/crates/office/skills/a3s-use-office/references/mcp.md @@ -29,6 +29,8 @@ Use its typed tools rather than passing shell command strings: - `office_save` persists a mutable session. - `office_close` refuses unsaved changes unless `discard=true` is explicit. - `office_list` reports sessions owned by this server process. +- `office_install_compat` prepares the optional pinned compatibility provider; + in Code it must pass parent confirmation before network access. Mutations remain unsaved until `office_save`. Do not discard a dirty session unless the user explicitly accepts losing its changes. Release the session as @@ -619,7 +621,10 @@ provider; other native Office tools do not require Browser or OfficeCLI. In an A3S Code `use` worker, use an available `mcp__use_office_compat__*` tool only when the native vocabulary lacks the -requested operation. In a CLI-only MCP host, `a3s use mcp serve office-compat` -starts the pinned OfficeCLI compatibility server; the legacy -`a3s use mcp serve office` alias remains supported. It is a separate standard -MCP target and is not the native session engine. +requested operation. If that surface is missing, call +`mcp__use_office__office_install_compat` through parent confirmation and wait +for the host to project the ready compatibility route. In a CLI-only MCP host, +`a3s use mcp serve office-compat` prepares and starts the pinned OfficeCLI +compatibility server; the legacy `a3s use mcp serve office` alias remains +supported. It is a separate standard MCP target and is not the native session +engine. diff --git a/docs/architecture.md b/docs/architecture.md index 4dc16a81..7e196cc6 100644 --- a/docs/architecture.md +++ b/docs/architecture.md @@ -15,15 +15,17 @@ compatibility backend until the native promotion gates in Search depends directly on the object-safe PageRenderer contract in a3s-use-browser. It never executes the CLI or requires a background service. -Provider selection is explicit. `DiscoveredChrome` is the default and never -downloads software. Only a `Managed*` provider or an explicit component install -authorizes a download. Managed downloads are restricted to approved HTTPS -hosts and redirects, bounded by size, hashed into an installation receipt, -staged outside the active version, and atomically activated. Lightpanda assets -must match the publisher SHA-256 exposed by GitHub Releases. Chrome for Testing's -current version feed does not publish an independent SHA-256 value, so its -receipt records HTTPS provenance and locally observed hashes without claiming -publisher checksum verification. +Provider selection is typed. `DiscoveredChrome` remains the non-installing +default for embedded callers such as Search. A validated direct A3S Browser +launch selects first-use preparation, while a Code worker requests the same +bounded installer through parent confirmation. Both reuse system Chrome and +the shared A3S-managed cache before downloading. Managed downloads are +restricted to approved HTTPS hosts and redirects, bounded by size, hashed into +an installation receipt, staged outside the active version, and atomically +activated. Lightpanda assets must match the publisher SHA-256 exposed by GitHub +Releases. Chrome for Testing's current version feed does not publish an +independent SHA-256 value, so its receipt records HTTPS provenance and locally +observed hashes without claiming publisher checksum verification. ## Native extension surfaces @@ -109,8 +111,10 @@ compatibility provider. For resident hosts, `use/office` targets the built-in OfficeCLI provider is projected separately as `use/office-compat`, targeting the standard compatibility server without carrying the native Skill. The `use/ocr` route targets `ocr-native`; model readiness remains visible. -Read-only discovery never installs models, while direct extraction uses the -first-use policy and the MCP worker uses a separately confirmed install tool. +Read-only discovery never installs providers. Direct Browser launch, OfficeCLI +compatibility execution, and OCR extraction use the first-use policy. Their MCP +workers use separately annotated installers that require parent confirmation; +native Office needs no provider installation. The projection contains content-bound Skill references and an MCP launch target, never executable extension code or a generic action payload. Consumers still @@ -440,8 +444,9 @@ Each invocation accepts argv and returns one versioned JSON document plus an exit status. This is CLI automation, not JSON-RPC. In 0.1.x, managed Office installation means the reviewed OfficeCLI compatibility -release. It is fetched only by an explicit install or repair command, restricted -to approved HTTPS hosts, bounded by size, and checked against the publisher's +release. It is fetched by explicit preparation, the first real compatibility +CLI command, or the confirmed native MCP installer. Downloads are restricted to +approved HTTPS hosts, bounded by size, and checked against the publisher's SHA-256 before atomic activation. Compatibility execution sets `OFFICECLI_SKIP_UPDATE=1`; A3S upgrades are explicit component operations. diff --git a/src/browser_cli.rs b/src/browser_cli.rs index 0ff0beb6..a581ad84 100644 --- a/src/browser_cli.rs +++ b/src/browser_cli.rs @@ -20,6 +20,7 @@ pub(crate) async fn run(args: &[String]) -> UseResult { return Ok(render_help()); } let options = RenderOptions::parse(&args[1..])?; + crate::first_use::ensure_browser_ready().await?; let pool = Arc::new(BrowserPool::new(BrowserPoolConfig::default())); let result = render_with(Arc::clone(&pool), options).await; pool.shutdown().await; diff --git a/src/cli.rs b/src/cli.rs index d869434a..1f41f5bb 100644 --- a/src/cli.rs +++ b/src/cli.rs @@ -684,6 +684,8 @@ async fn browser(args: &[String]) -> UseResult { async fn office(args: &[String]) -> UseResult { match args.first().map(String::as_str) { None | Some("doctor") => doctor(Some("office")), + Some("-h" | "--help" | "help") => Ok(office_help()), + Some("-V" | "--version" | "version") => Ok(version()), Some("skills") => { #[cfg(feature = "office")] return crate::office_skills::run(&args[1..]).await; @@ -705,6 +707,7 @@ async fn office(args: &[String]) -> UseResult { Some(_) => { #[cfg(feature = "office")] { + crate::first_use::ensure_office_compatibility_ready().await?; let exit_code = a3s_use_office::delegate_native(args).await?; Ok(CommandOutput::delegated(exit_code)) } @@ -717,6 +720,18 @@ async fn office(args: &[String]) -> UseResult { } } +fn office_help() -> CommandOutput { + let usage = concat!( + "usage:\n", + " a3s-use office doctor [--json]\n", + " a3s-use office skills list|get|path [args] [--json]\n", + " a3s-use office native [args] [--json]\n", + " a3s-use office [args]\n\n", + "Native Office is built in. The optional OfficeCLI compatibility provider is prepared on its first real command when policy allows." + ); + CommandOutput::success(usage, serde_json::json!({ "usage": usage })) +} + async fn extension(args: &[String]) -> UseResult { match args.first().map(String::as_str) { None | Some("list") => extension_list().await, @@ -804,6 +819,7 @@ async fn mcp(args: &[String]) -> UseResult { } #[cfg(feature = "office")] { + crate::first_use::ensure_office_compatibility_ready().await?; let exit_code = a3s_use_office::delegate_native(&["mcp".to_string()]).await?; Ok(CommandOutput::delegated(exit_code)) diff --git a/src/first_use.rs b/src/first_use.rs new file mode 100644 index 00000000..7f93d708 --- /dev/null +++ b/src/first_use.rs @@ -0,0 +1,145 @@ +use a3s_use_core::{FirstUseInstallPolicy, UseError, UseResult}; + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +enum AutoInstallAction { + Ready, + Install, +} + +fn automatic_install_action( + domain: &str, + asset: &str, + available: bool, + explicit_invalid: bool, + policy: FirstUseInstallPolicy, +) -> UseResult { + if explicit_invalid { + return Err(UseError::new( + format!("use.{domain}.explicit_provider_invalid"), + format!("The explicit {domain} provider is not usable."), + ) + .with_suggestion(format!( + "Fix or unset the explicit {domain} provider before retrying." + ))); + } + if available { + return Ok(AutoInstallAction::Ready); + } + if let Some(block) = policy.blocked_by() { + return Err(UseError::new( + format!("use.{domain}.auto_install_disabled"), + format!( + "{asset} is not ready and first-use installation is disabled by {}.", + block.reason() + ), + ) + .with_suggestion(format!( + "Enable first-use installation or prepare {asset} explicitly while online." + )) + .with_detail("reason", block.reason())); + } + Ok(AutoInstallAction::Install) +} + +#[cfg(feature = "browser")] +pub(crate) async fn ensure_browser_ready() -> UseResult { + use a3s_use_browser::{BrowserInstallSource, ManagedBrowser}; + + let status = a3s_use_browser::browser_status(ManagedBrowser::Chrome); + let explicit_configured = explicit_environment_value(&["A3S_BROWSER_EXECUTABLE", "CHROME"]); + let explicit_invalid = explicit_configured + && !(status.available && status.source == BrowserInstallSource::Environment); + match automatic_install_action( + "browser", + "the shared A3S Use Browser runtime", + status.available, + explicit_invalid, + FirstUseInstallPolicy::from_env()?, + )? { + AutoInstallAction::Ready => Ok(status), + AutoInstallAction::Install => { + a3s_use_browser::install_browser(ManagedBrowser::Chrome).await + } + } +} + +#[cfg(feature = "office")] +pub(crate) async fn ensure_office_compatibility_ready( +) -> UseResult { + use a3s_use_office::OfficeInstallSource; + + let status = a3s_use_office::office_status(); + let explicit_configured = explicit_environment_value(&["A3S_OFFICECLI_EXECUTABLE"]); + let explicit_invalid = explicit_configured + && !(status.available && status.source == OfficeInstallSource::Environment); + match automatic_install_action( + "office", + "the optional OfficeCLI compatibility provider", + status.available, + explicit_invalid, + FirstUseInstallPolicy::from_env()?, + )? { + AutoInstallAction::Ready => Ok(status), + AutoInstallAction::Install => a3s_use_office::install_office_cli(false).await, + } +} + +#[cfg(any(feature = "browser", feature = "office"))] +fn explicit_environment_value(names: &[&str]) -> bool { + names + .iter() + .any(|name| std::env::var_os(name).is_some_and(|value| !value.is_empty())) +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn ready_provider_never_installs_or_fails_policy() { + let action = automatic_install_action( + "browser", + "Browser", + true, + false, + FirstUseInstallPolicy::new(true, true), + ) + .unwrap(); + assert_eq!(action, AutoInstallAction::Ready); + } + + #[test] + fn missing_provider_installs_when_policy_allows_it() { + let action = automatic_install_action( + "office", + "OfficeCLI", + false, + false, + FirstUseInstallPolicy::new(false, false), + ) + .unwrap(); + assert_eq!(action, AutoInstallAction::Install); + } + + #[test] + fn policy_and_explicit_provider_boundaries_are_typed() { + let explicit = automatic_install_action( + "browser", + "Browser", + false, + true, + FirstUseInstallPolicy::new(false, false), + ) + .unwrap_err(); + assert_eq!(explicit.code, "use.browser.explicit_provider_invalid"); + + for policy in [ + FirstUseInstallPolicy::new(true, false), + FirstUseInstallPolicy::new(false, true), + ] { + let error = + automatic_install_action("office", "OfficeCLI", false, false, policy).unwrap_err(); + assert_eq!(error.code, "use.office.auto_install_disabled"); + } + } +} diff --git a/src/lib.rs b/src/lib.rs index fe1d9d97..c4d71b43 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -10,6 +10,7 @@ mod capability_registry; pub mod cli; mod component_route; mod extension_cli; +mod first_use; #[cfg(feature = "ocr")] mod ocr_builtin; diff --git a/src/mcp/office.rs b/src/mcp/office.rs index 3305ff59..bfc6c201 100644 --- a/src/mcp/office.rs +++ b/src/mcp/office.rs @@ -84,6 +84,25 @@ impl NativeOfficeMcpServer { #[tool_router] impl NativeOfficeMcpServer { + #[tool( + name = "office_install_compat", + description = "Install or repair the optional pinned OfficeCLI compatibility provider through the bounded A3S component lifecycle", + annotations( + read_only_hint = false, + destructive_hint = false, + idempotent_hint = true, + open_world_hint = true + ) + )] + async fn office_install_compat(&self) -> Result { + let result = async { + let status = crate::first_use::ensure_office_compatibility_ready().await?; + serde_json::to_value(status).map_err(output_encoding_error) + } + .await; + Ok(tool_result(result)) + } + #[tool( name = "office_validate", description = "Validate and identify one local OOXML document without opening a session", @@ -536,7 +555,7 @@ impl ServerHandler for NativeOfficeMcpServer { website_url: Some("https://github.com/A3S-Lab/Use".to_string()), }, instructions: Some( - "This explicit preview server edits OOXML in process and never installs or starts OfficeCLI, Microsoft Office, or LibreOffice. Create or open a session first. Mutations remain in memory until office_save; office_close refuses unsaved changes unless discard=true. The separate `mcp serve office` target remains the OfficeCLI compatibility server until native promotion gates pass." + "Use the built-in native Office tools first; they never require OfficeCLI, Microsoft Office, or LibreOffice. If a requested operation is outside the native surface, request office_install_compat through the host confirmation path and use the separately projected Office compatibility route after it becomes ready. Create or open a native session first. Mutations remain in memory until office_save; office_close refuses unsaved changes unless discard=true." .to_string(), ), ..Default::default() diff --git a/src/mcp/office/tests.rs b/src/mcp/office/tests.rs index d6e5bc6d..7100d850 100644 --- a/src/mcp/office/tests.rs +++ b/src/mcp/office/tests.rs @@ -11,7 +11,7 @@ use a3s_use_office::{ }; #[test] -fn native_office_server_exposes_only_bounded_typed_tools() { +fn native_office_server_exposes_bounded_tools_and_confirmed_compat_install() { let server = NativeOfficeMcpServer::new(); let tools = server.tool_router.list_all(); let mut names: Vec<&str> = tools @@ -26,6 +26,7 @@ fn native_office_server_exposes_only_bounded_typed_tools() { "office_close", "office_create", "office_get", + "office_install_compat", "office_list", "office_merge_template", "office_open", @@ -60,6 +61,18 @@ fn native_office_server_exposes_only_bounded_typed_tools() { Some(false) ); assert_eq!(annotations("office_save").destructive_hint, Some(true)); + assert_eq!( + annotations("office_install_compat").read_only_hint, + Some(false) + ); + assert_eq!( + annotations("office_install_compat").idempotent_hint, + Some(true) + ); + assert_eq!( + annotations("office_install_compat").open_world_hint, + Some(true) + ); } #[test] diff --git a/tests/cli.rs b/tests/cli.rs index 4c1034f3..9eae45b9 100644 --- a/tests/cli.rs +++ b/tests/cli.rs @@ -1835,6 +1835,50 @@ fn office_install_reuses_an_explicit_provider_without_downloading() { assert!(!temp.path().join("managed/1.0.136").exists()); } +#[cfg(feature = "office")] +#[test] +fn office_compatibility_first_use_honors_no_auto_install() { + let temp = tempfile::tempdir().unwrap(); + let output = Command::new(binary()) + .args(["office", "document", "inspect", "fixture.docx", "--json"]) + .env("A3S_USE_OFFICE_HOME", temp.path().join("managed")) + .env("A3S_NO_AUTO_INSTALL", "1") + .env("PATH", temp.path()) + .output() + .unwrap(); + + assert_eq!(output.status.code(), Some(1), "{output:?}"); + let value: serde_json::Value = serde_json::from_slice(&output.stdout).unwrap(); + assert_eq!(value["ok"], false); + assert_eq!(value["error"]["code"], "use.office.auto_install_disabled"); + assert_eq!(value["error"]["details"]["reason"], "A3S_NO_AUTO_INSTALL"); + assert!(!temp.path().join("managed/1.0.136").exists()); +} + +#[cfg(feature = "office")] +#[test] +fn office_help_never_triggers_compatibility_installation() { + let temp = tempfile::tempdir().unwrap(); + let output = Command::new(binary()) + .args(["office", "--help", "--json"]) + .env( + "A3S_OFFICECLI_EXECUTABLE", + temp.path().join("must-not-exist"), + ) + .env("A3S_USE_OFFICE_HOME", temp.path().join("managed")) + .output() + .unwrap(); + + assert!(output.status.success(), "{output:?}"); + let value: serde_json::Value = serde_json::from_slice(&output.stdout).unwrap(); + assert_eq!(value["ok"], true); + assert!(value["data"]["usage"] + .as_str() + .unwrap() + .contains("Native Office is built in")); + assert!(!temp.path().join("managed/1.0.136").exists()); +} + #[cfg(all(unix, feature = "office"))] #[test] fn office_mcp_target_delegates_to_officeclis_standard_server() { @@ -1949,7 +1993,15 @@ async fn native_office_mcp_is_standard_typed_and_independent_of_officecli() { ) .await; let tools = tools["result"]["tools"].as_array().unwrap(); - assert_eq!(tools.len(), 12); + assert_eq!(tools.len(), 13); + let install_compat = tools + .iter() + .find(|tool| tool["name"] == "office_install_compat") + .unwrap(); + assert_eq!(install_compat["annotations"]["readOnlyHint"], false); + assert_eq!(install_compat["annotations"]["destructiveHint"], false); + assert_eq!(install_compat["annotations"]["idempotentHint"], true); + assert_eq!(install_compat["annotations"]["openWorldHint"], true); let apply = tools .iter() .find(|tool| tool["name"] == "office_apply_batch") @@ -2196,6 +2248,27 @@ fn built_in_ocr_projects_the_canonical_code_route_and_skill() { assert_eq!(digest.len(), 64); } +#[cfg(feature = "browser")] +#[test] +fn browser_render_never_replaces_an_invalid_explicit_provider() { + let temp = tempfile::tempdir().unwrap(); + let output = Command::new(binary()) + .args(["browser", "render", "https://example.com", "--json"]) + .env("A3S_BROWSER_EXECUTABLE", temp.path().join("must-not-exist")) + .env("A3S_USE_BROWSER_HOME", temp.path().join("managed")) + .output() + .unwrap(); + + assert_eq!(output.status.code(), Some(1), "{output:?}"); + let value: serde_json::Value = serde_json::from_slice(&output.stdout).unwrap(); + assert_eq!(value["ok"], false); + assert_eq!( + value["error"]["code"], + "use.browser.explicit_provider_invalid" + ); + assert!(!temp.path().join("managed/chrome").exists()); +} + #[cfg(feature = "ocr")] #[test] fn ocr_extract_honors_the_no_auto_install_boundary_for_a_valid_image() { From d5c913db5596f92c17b99cf28b82d44376813f50 Mon Sep 17 00:00:00 2001 From: RoyLin Date: Sun, 19 Jul 2026 13:30:17 +0800 Subject: [PATCH 8/9] refactor(ocr): remove obsolete provider implementation --- crates/ocr/src/provider.rs | 393 ------------------------------------- 1 file changed, 393 deletions(-) delete mode 100644 crates/ocr/src/provider.rs diff --git a/crates/ocr/src/provider.rs b/crates/ocr/src/provider.rs deleted file mode 100644 index eb8880e0..00000000 --- a/crates/ocr/src/provider.rs +++ /dev/null @@ -1,393 +0,0 @@ -use std::env; -use std::path::{Path, PathBuf}; -use std::time::Duration; - -use a3s_use_core::{Readiness, UseError, UseResult}; -use url::Url; - -use crate::{OcrDiagnostic, OcrProviderKind}; - -const DEFAULT_TIMEOUT: Duration = Duration::from_secs(60); -const DEFAULT_VISION_BASE_URL: &str = "https://api.openai.com/v1/"; - -#[derive(Debug, Clone)] -pub(crate) enum Provider { - Tesseract { - executable: PathBuf, - timeout: Duration, - }, - Vision { - endpoint: Url, - api_key: Option, - model: String, - timeout: Duration, - }, -} - -impl Provider { - pub(crate) fn kind(&self) -> OcrProviderKind { - match self { - Self::Tesseract { .. } => OcrProviderKind::Tesseract, - Self::Vision { .. } => OcrProviderKind::Vision, - } - } - - pub(crate) fn diagnostic(&self) -> OcrDiagnostic { - match self { - Self::Tesseract { executable, .. } => OcrDiagnostic { - readiness: Readiness::Ready, - provider: Some(OcrProviderKind::Tesseract), - executable: Some(executable.clone()), - endpoint: None, - model: None, - sends_source_off_device: false, - message: "The local Tesseract OCR provider is ready.".to_string(), - suggestions: Vec::new(), - }, - Self::Vision { - endpoint, model, .. - } => OcrDiagnostic { - readiness: Readiness::Ready, - provider: Some(OcrProviderKind::Vision), - executable: None, - endpoint: Some(redacted_endpoint(endpoint)), - model: Some(model.clone()), - sends_source_off_device: !is_loopback(endpoint), - message: "The explicitly configured vision OCR provider is ready.".to_string(), - suggestions: Vec::new(), - }, - } - } -} - -#[derive(Debug, Clone)] -pub(crate) struct ProviderConfig { - requested: OcrProviderKind, - tesseract: Option, - vision: Option, - timeout: Duration, -} - -#[derive(Debug, Clone)] -struct VisionConfig { - endpoint: Url, - api_key: Option, - model: String, -} - -impl ProviderConfig { - pub(crate) fn from_env() -> UseResult { - let requested = match env::var("A3S_OCR_PROVIDER") - .unwrap_or_else(|_| "auto".to_string()) - .trim() - { - "" | "auto" => OcrProviderKind::Auto, - "tesseract" => OcrProviderKind::Tesseract, - "vision" => OcrProviderKind::Vision, - value => { - return Err(UseError::new( - "use.ocr.provider_invalid", - format!("Unknown OCR provider '{value}'; expected auto, tesseract, or vision."), - )) - } - }; - - let timeout = timeout_from_env()?; - let tesseract = env::var_os("A3S_OCR_TESSERACT_EXECUTABLE") - .filter(|value| !value.is_empty()) - .map(PathBuf::from) - .or_else(|| find_on_path("tesseract")); - let vision = vision_config_from_env()?; - Ok(Self { - requested, - tesseract, - vision, - timeout, - }) - } - - #[cfg(all(test, unix))] - pub(crate) fn tesseract(executable: PathBuf) -> Self { - Self { - requested: OcrProviderKind::Tesseract, - tesseract: Some(executable), - vision: None, - timeout: DEFAULT_TIMEOUT, - } - } - - pub(crate) fn diagnostic(&self) -> OcrDiagnostic { - match self.resolve(self.requested) { - Ok(provider) => provider.diagnostic(), - Err(error) => OcrDiagnostic { - readiness: Readiness::Missing, - provider: match self.requested { - OcrProviderKind::Auto => None, - provider => Some(provider), - }, - executable: self.tesseract.clone(), - endpoint: self - .vision - .as_ref() - .map(|vision| redacted_endpoint(&vision.endpoint)), - model: self.vision.as_ref().map(|vision| vision.model.clone()), - sends_source_off_device: self - .vision - .as_ref() - .is_some_and(|vision| !is_loopback(&vision.endpoint)), - message: error.message, - suggestions: error.suggestion.into_iter().collect(), - }, - } - } - - pub(crate) fn resolve(&self, requested: OcrProviderKind) -> UseResult { - let requested = if requested == OcrProviderKind::Auto { - self.requested - } else { - requested - }; - match requested { - OcrProviderKind::Auto => { - if let Some(executable) = &self.tesseract { - return tesseract_provider(executable, self.timeout); - } - if let Some(vision) = &self.vision { - return Ok(vision_provider(vision, self.timeout)); - } - Err(missing_provider()) - } - OcrProviderKind::Tesseract => self - .tesseract - .as_ref() - .ok_or_else(missing_tesseract) - .and_then(|path| tesseract_provider(path, self.timeout)), - OcrProviderKind::Vision => self - .vision - .as_ref() - .map(|vision| vision_provider(vision, self.timeout)) - .ok_or_else(missing_vision), - } - } -} - -fn tesseract_provider(path: &Path, timeout: Duration) -> UseResult { - let path = std::fs::canonicalize(path).map_err(|error| { - UseError::new( - "use.ocr.provider_missing", - format!( - "Configured Tesseract executable '{}' is not readable: {error}", - path.display() - ), - ) - .with_suggestion( - "Install Tesseract explicitly or configure the vision provider; A3S Use will not install an OCR provider automatically.", - ) - })?; - let metadata = std::fs::metadata(&path).map_err(|error| { - UseError::new( - "use.ocr.provider_missing", - format!( - "Configured Tesseract executable '{}' is not readable: {error}", - path.display() - ), - ) - })?; - if !metadata.is_file() { - return Err(UseError::new( - "use.ocr.provider_invalid", - format!( - "Configured Tesseract path '{}' is not a regular file.", - path.display() - ), - )); - } - #[cfg(unix)] - { - use std::os::unix::fs::PermissionsExt; - if metadata.permissions().mode() & 0o111 == 0 { - return Err(UseError::new( - "use.ocr.provider_invalid", - format!( - "Configured Tesseract path '{}' is not executable.", - path.display() - ), - )); - } - } - Ok(Provider::Tesseract { - executable: path, - timeout, - }) -} - -fn vision_provider(config: &VisionConfig, timeout: Duration) -> Provider { - Provider::Vision { - endpoint: config.endpoint.clone(), - api_key: config.api_key.clone(), - model: config.model.clone(), - timeout, - } -} - -fn vision_config_from_env() -> UseResult> { - let model = env::var("A3S_OCR_VISION_MODEL") - .ok() - .map(|value| value.trim().to_string()) - .filter(|value| !value.is_empty()); - let base_url = env::var("A3S_OCR_VISION_BASE_URL") - .ok() - .map(|value| value.trim().to_string()) - .filter(|value| !value.is_empty()); - let api_key = env::var("A3S_OCR_VISION_API_KEY") - .ok() - .map(|value| value.trim().to_string()) - .filter(|value| !value.is_empty()); - - if model.is_none() && base_url.is_none() && api_key.is_none() { - return Ok(None); - } - let model = model.ok_or_else(|| { - UseError::new( - "use.ocr.vision_config_invalid", - "A3S_OCR_VISION_MODEL is required when the vision OCR provider is configured.", - ) - })?; - let mut base = base_url.unwrap_or_else(|| DEFAULT_VISION_BASE_URL.to_string()); - if !base.ends_with('/') { - base.push('/'); - } - let base = Url::parse(&base).map_err(|error| { - UseError::new( - "use.ocr.vision_config_invalid", - format!("A3S_OCR_VISION_BASE_URL is invalid: {error}"), - ) - })?; - validate_endpoint(&base, api_key.as_deref())?; - let endpoint = base.join("chat/completions").map_err(|error| { - UseError::new( - "use.ocr.vision_config_invalid", - format!("Failed to resolve the vision OCR endpoint: {error}"), - ) - })?; - Ok(Some(VisionConfig { - endpoint, - api_key, - model, - })) -} - -fn validate_endpoint(endpoint: &Url, api_key: Option<&str>) -> UseResult<()> { - if !endpoint.username().is_empty() || endpoint.password().is_some() { - return Err(UseError::new( - "use.ocr.vision_config_invalid", - "The vision OCR endpoint must not contain embedded credentials.", - )); - } - if endpoint.scheme() != "https" && !(endpoint.scheme() == "http" && is_loopback(endpoint)) { - return Err(UseError::new( - "use.ocr.vision_config_invalid", - "The vision OCR endpoint must use HTTPS; loopback HTTP is allowed for local providers.", - )); - } - if !is_loopback(endpoint) && api_key.is_none() { - return Err(UseError::new( - "use.ocr.vision_config_invalid", - "A3S_OCR_VISION_API_KEY is required for a non-loopback vision endpoint.", - )); - } - Ok(()) -} - -fn timeout_from_env() -> UseResult { - let Some(value) = env::var("A3S_OCR_TIMEOUT_MS").ok() else { - return Ok(DEFAULT_TIMEOUT); - }; - let millis = value.parse::().map_err(|_| { - UseError::new( - "use.ocr.timeout_invalid", - "A3S_OCR_TIMEOUT_MS must be an integer from 1 through 300000.", - ) - })?; - if !(1..=300_000).contains(&millis) { - return Err(UseError::new( - "use.ocr.timeout_invalid", - "A3S_OCR_TIMEOUT_MS must be an integer from 1 through 300000.", - )); - } - Ok(Duration::from_millis(millis)) -} - -fn find_on_path(name: &str) -> Option { - let path = env::var_os("PATH")?; - env::split_paths(&path) - .map(|directory| directory.join(executable_name(name))) - .find(|candidate| candidate.is_file()) -} - -fn executable_name(name: &str) -> String { - if cfg!(windows) { - format!("{name}.exe") - } else { - name.to_string() - } -} - -fn missing_provider() -> UseError { - UseError::new( - "use.ocr.provider_missing", - "No OCR provider is configured or discoverable.", - ) - .with_suggestion( - "Install Tesseract explicitly, set A3S_OCR_TESSERACT_EXECUTABLE, or configure A3S_OCR_VISION_MODEL, A3S_OCR_VISION_BASE_URL, and A3S_OCR_VISION_API_KEY.", - ) -} - -fn missing_tesseract() -> UseError { - UseError::new( - "use.ocr.provider_missing", - "The Tesseract OCR provider is not installed or configured.", - ) - .with_suggestion( - "Install Tesseract explicitly or set A3S_OCR_TESSERACT_EXECUTABLE; A3S Use will not install it automatically.", - ) -} - -fn missing_vision() -> UseError { - UseError::new( - "use.ocr.provider_missing", - "The vision OCR provider is not configured.", - ) - .with_suggestion( - "Set A3S_OCR_VISION_MODEL and an approved HTTPS endpoint/API key before sending source images to a vision provider.", - ) -} - -fn is_loopback(url: &Url) -> bool { - matches!(url.host_str(), Some("localhost" | "127.0.0.1" | "::1")) -} - -fn redacted_endpoint(url: &Url) -> String { - let mut redacted = url.clone(); - redacted.set_query(None); - redacted.set_fragment(None); - redacted.to_string() -} - -#[cfg(test)] -mod tests { - use super::*; - - #[test] - fn rejects_insecure_remote_vision_endpoint() { - let endpoint = Url::parse("http://ocr.example.com/v1/").unwrap(); - let error = validate_endpoint(&endpoint, Some("secret")).unwrap_err(); - assert_eq!(error.code, "use.ocr.vision_config_invalid"); - } - - #[test] - fn permits_loopback_http_without_an_api_key() { - let endpoint = Url::parse("http://127.0.0.1:8080/v1/").unwrap(); - validate_endpoint(&endpoint, None).unwrap(); - } -} From 9280837c79b05a59365d3d5c85db9ab6b04734c6 Mon Sep 17 00:00:00 2001 From: RoyLin Date: Sun, 19 Jul 2026 13:54:40 +0800 Subject: [PATCH 9/9] release: prepare a3s-use 0.1.2 --- .github/workflows/release.yml | 4 ++- Cargo.lock | 26 ++++++++++++-------- Cargo.toml | 12 ++++----- README.md | 2 +- crates/browser-driver/Cargo.toml | 4 +-- crates/browser/Cargo.toml | 2 +- crates/extension/Cargo.toml | 4 +-- crates/ocr/Cargo.toml | 2 +- crates/office/Cargo.toml | 2 +- crates/science/Cargo.toml | 4 +-- crates/science/package/a3s-use-extension.acl | 2 +- 11 files changed, 36 insertions(+), 28 deletions(-) diff --git a/.github/workflows/release.yml b/.github/workflows/release.yml index 899d5886..e201558c 100644 --- a/.github/workflows/release.yml +++ b/.github/workflows/release.yml @@ -316,7 +316,7 @@ jobs: ref: ${{ github.event_name == 'workflow_dispatch' && inputs.release_tag || github.ref }} - uses: dtolnay/rust-toolchain@stable - uses: Swatinem/rust-cache@v2 - - name: Publish Core, OCR, then Browser + - name: Publish Core, Extension, OCR, then Browser env: CARGO_REGISTRY_TOKEN: ${{ secrets.CARGO_TOKEN }} VERSION: ${{ needs.validate.outputs.version }} @@ -355,6 +355,8 @@ jobs: publish_once a3s-use-core wait_until_visible a3s-use-core + publish_once a3s-use-extension + wait_until_visible a3s-use-extension publish_once a3s-use-ocr wait_until_visible a3s-use-ocr publish_once a3s-use-browser diff --git a/Cargo.lock b/Cargo.lock index 7ef15ef8..aa7c726e 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -7,9 +7,15 @@ name = "a3s-acl" version = "0.2.1" source = "git+https://github.com/A3S-Lab/ACL?rev=6e2a6469edc0f4c61b1e588d0ace873aaf15ce22#6e2a6469edc0f4c61b1e588d0ace873aaf15ce22" +[[package]] +name = "a3s-acl" +version = "0.2.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d2dc4eb3b0dd1b11efa0ad9bf397c97fd1d16e3eb4f9ca43872df069560d0b69" + [[package]] name = "a3s-use" -version = "0.1.1" +version = "0.1.2" dependencies = [ "a3s-use-browser", "a3s-use-core", @@ -44,7 +50,7 @@ dependencies = [ [[package]] name = "a3s-use-browser" -version = "0.1.1" +version = "0.1.2" dependencies = [ "a3s-use-core", "async-trait", @@ -65,9 +71,9 @@ dependencies = [ [[package]] name = "a3s-use-browser-driver" -version = "0.1.1" +version = "0.1.2" dependencies = [ - "a3s-acl", + "a3s-acl 0.2.1", "a3s-use-core", "aes-gcm", "async-trait", @@ -99,7 +105,7 @@ dependencies = [ [[package]] name = "a3s-use-core" -version = "0.1.1" +version = "0.1.2" dependencies = [ "serde", "serde_json", @@ -108,9 +114,9 @@ dependencies = [ [[package]] name = "a3s-use-extension" -version = "0.1.1" +version = "0.1.2" dependencies = [ - "a3s-acl", + "a3s-acl 0.2.2", "a3s-use-core", "flate2", "fs2", @@ -131,7 +137,7 @@ dependencies = [ [[package]] name = "a3s-use-ocr" -version = "0.1.1" +version = "0.1.2" dependencies = [ "a3s-use-core", "clap", @@ -155,7 +161,7 @@ dependencies = [ [[package]] name = "a3s-use-office" -version = "0.1.1" +version = "0.1.2" dependencies = [ "a3s-use-core", "async-trait", @@ -175,7 +181,7 @@ dependencies = [ [[package]] name = "a3s-use-science" -version = "0.1.1" +version = "0.1.2" dependencies = [ "a3s-use-core", "a3s-use-extension", diff --git a/Cargo.toml b/Cargo.toml index efef3f83..d4710d39 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -11,7 +11,7 @@ members = [ resolver = "2" [workspace.package] -version = "0.1.1" +version = "0.1.2" edition = "2021" license = "MIT" repository = "https://github.com/A3S-Lab/Use" @@ -92,11 +92,11 @@ mcp = [ lightpanda = ["browser", "a3s-use-browser/lightpanda"] [dependencies] -a3s-use-core = { version = "0.1.1", path = "crates/core" } -a3s-use-browser = { version = "0.1.1", path = "crates/browser", optional = true } -a3s-use-office = { version = "0.1.1", path = "crates/office", optional = true } -a3s-use-ocr = { version = "0.1.1", path = "crates/ocr", optional = true } -a3s-use-extension = { version = "0.1.1", path = "crates/extension", optional = true } +a3s-use-core = { version = "0.1.2", path = "crates/core" } +a3s-use-browser = { version = "0.1.2", path = "crates/browser", optional = true } +a3s-use-office = { version = "0.1.2", path = "crates/office", optional = true } +a3s-use-ocr = { version = "0.1.2", path = "crates/ocr", optional = true } +a3s-use-extension = { version = "0.1.2", path = "crates/extension", optional = true } anyhow.workspace = true axum = { workspace = true, optional = true } base64 = { workspace = true, optional = true } diff --git a/README.md b/README.md index 5b0cc735..20fa1e19 100644 --- a/README.md +++ b/README.md @@ -244,7 +244,7 @@ not the facade binary: ```toml [dependencies] -a3s-use-browser = "0.1.1" +a3s-use-browser = "0.1.2" tokio = { version = "1", features = ["macros", "rt-multi-thread"] } url = "2" ``` diff --git a/crates/browser-driver/Cargo.toml b/crates/browser-driver/Cargo.toml index 8f0828ca..3e837773 100644 --- a/crates/browser-driver/Cargo.toml +++ b/crates/browser-driver/Cargo.toml @@ -1,6 +1,6 @@ [package] name = "a3s-use-browser-driver" -version = "0.1.1" +version = "0.1.2" edition = "2021" description = "Full browser automation driver embedded in A3S Use" license = "Apache-2.0" @@ -13,7 +13,7 @@ name = "a3s-use-browser-driver" path = "src/main.rs" [dependencies] -a3s-use-core = { version = "0.1.1", path = "../core" } +a3s-use-core = { version = "0.1.2", path = "../core" } a3s-acl = { git = "https://github.com/A3S-Lab/ACL", rev = "6e2a6469edc0f4c61b1e588d0ace873aaf15ce22" } serde = { version = "1.0", features = ["derive"] } serde_json = "1.0" diff --git a/crates/browser/Cargo.toml b/crates/browser/Cargo.toml index 88f05d6d..4147e314 100644 --- a/crates/browser/Cargo.toml +++ b/crates/browser/Cargo.toml @@ -22,7 +22,7 @@ chrome = [ lightpanda = ["chrome"] [dependencies] -a3s-use-core = { version = "0.1.1", path = "../core" } +a3s-use-core = { version = "0.1.2", path = "../core" } async-trait.workspace = true chromiumoxide = { version = "0.7", features = ["tokio-runtime"], optional = true } fs2 = { workspace = true, optional = true } diff --git a/crates/extension/Cargo.toml b/crates/extension/Cargo.toml index d1ea839f..5a70b363 100644 --- a/crates/extension/Cargo.toml +++ b/crates/extension/Cargo.toml @@ -9,8 +9,8 @@ rust-version.workspace = true description = "ACL manifest and native surface contracts for A3S Use extensions" [dependencies] -a3s-acl = { git = "https://github.com/A3S-Lab/ACL", rev = "6e2a6469edc0f4c61b1e588d0ace873aaf15ce22" } -a3s-use-core = { version = "0.1.1", path = "../core" } +a3s-acl = "=0.2.2" +a3s-use-core = { version = "0.1.2", path = "../core" } fs2.workspace = true flate2.workspace = true reqwest.workspace = true diff --git a/crates/ocr/Cargo.toml b/crates/ocr/Cargo.toml index dac1f616..e9ba0dc7 100644 --- a/crates/ocr/Cargo.toml +++ b/crates/ocr/Cargo.toml @@ -17,7 +17,7 @@ name = "a3s-use-ocr" path = "src/main.rs" [dependencies] -a3s-use-core = { version = "0.1.1", path = "../core" } +a3s-use-core = { version = "0.1.2", path = "../core" } clap.workspace = true clipper2.workspace = true fs2.workspace = true diff --git a/crates/office/Cargo.toml b/crates/office/Cargo.toml index e848dfb6..478c69ff 100644 --- a/crates/office/Cargo.toml +++ b/crates/office/Cargo.toml @@ -9,7 +9,7 @@ rust-version.workspace = true description = "Native OOXML operations and temporary OfficeCLI compatibility for A3S Use" [dependencies] -a3s-use-core = { version = "0.1.1", path = "../core" } +a3s-use-core = { version = "0.1.2", path = "../core" } async-trait.workspace = true base64.workspace = true fs2.workspace = true diff --git a/crates/science/Cargo.toml b/crates/science/Cargo.toml index 6e99d89d..f55c7ddf 100644 --- a/crates/science/Cargo.toml +++ b/crates/science/Cargo.toml @@ -17,7 +17,7 @@ name = "a3s-use-science" path = "src/main.rs" [dependencies] -a3s-use-core = { version = "0.1.1", path = "../core" } +a3s-use-core = { version = "0.1.2", path = "../core" } clap.workspace = true reqwest = { workspace = true, features = ["json"] } rmcp.workspace = true @@ -28,6 +28,6 @@ tokio.workspace = true url.workspace = true [dev-dependencies] -a3s-use-extension = { version = "0.1.1", path = "../extension" } +a3s-use-extension = { version = "0.1.2", path = "../extension" } axum.workspace = true tempfile.workspace = true diff --git a/crates/science/package/a3s-use-extension.acl b/crates/science/package/a3s-use-extension.acl index 97a5fe49..917c507e 100644 --- a/crates/science/package/a3s-use-extension.acl +++ b/crates/science/package/a3s-use-extension.acl @@ -1,6 +1,6 @@ extension "a3s/science" { schema_version = 1 - version = "0.1.1" + version = "0.1.2" route = "science" actions = ["read"]