diff --git a/.github/workflows/release.yml b/.github/workflows/release.yml index dfee133e..e201558c 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,16 @@ 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" + 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/" + 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" @@ -107,9 +111,15 @@ 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" + 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 +138,11 @@ 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}/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" @@ -143,6 +158,49 @@ 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"] == "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' + 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 +237,11 @@ 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/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", @@ -194,6 +257,36 @@ 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 + $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" + 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 +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 then Browser + - name: Publish Core, Extension, OCR, then Browser env: CARGO_REGISTRY_TOKEN: ${{ secrets.CARGO_TOKEN }} VERSION: ${{ needs.validate.outputs.version }} @@ -262,6 +355,10 @@ 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 release: diff --git a/Cargo.lock b/Cargo.lock index 14cb1eb3..aa7c726e 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -7,28 +7,39 @@ 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", "a3s-use-extension", + "a3s-use-ocr", "a3s-use-office", "anyhow", "async-trait", "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", @@ -39,7 +50,7 @@ dependencies = [ [[package]] name = "a3s-use-browser" -version = "0.1.1" +version = "0.1.2" dependencies = [ "a3s-use-core", "async-trait", @@ -60,9 +71,10 @@ 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", "base64", @@ -93,31 +105,63 @@ dependencies = [ [[package]] name = "a3s-use-core" -version = "0.1.1" +version = "0.1.2" dependencies = [ "serde", "serde_json", - "thiserror 2.0.18", + "thiserror 2.0.19", ] [[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", + "olpc-cjson", + "reqwest", + "ring", "semver", "serde", "serde_json", "sha2 0.10.9", + "tar", + "tempfile", + "tokio", + "tough", + "url", + "zip", +] + +[[package]] +name = "a3s-use-ocr" +version = "0.1.2" +dependencies = [ + "a3s-use-core", + "clap", + "clipper2", + "fs2", + "image", + "imageproc", + "ort", + "reqwest", + "rmcp", + "schemars", + "serde", + "serde_json", + "serde_yaml", + "sha2 0.10.9", + "tar", "tempfile", "tokio", + "url", ] [[package]] name = "a3s-use-office" -version = "0.1.1" +version = "0.1.2" dependencies = [ "a3s-use-core", "async-trait", @@ -135,6 +179,40 @@ dependencies = [ "zip", ] +[[package]] +name = "a3s-use-science" +version = "0.1.2" +dependencies = [ + "a3s-use-core", + "a3s-use-extension", + "axum", + "clap", + "reqwest", + "rmcp", + "schemars", + "serde", + "serde_json", + "tempfile", + "tokio", + "url", +] + +[[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" @@ -185,15 +263,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" @@ -264,9 +333,18 @@ dependencies = [ [[package]] name = "anyhow" -version = "1.0.103" +version = "1.0.104" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "330a5ed07fa54e4702c9d6c4174f74427fc0ef6e214bbd677ae50a5099946470" + +[[package]] +name = "approx" +version = "0.5.1" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "2a4385e2e34eb35d6b3efe798b9eb88096925d87726c0798709bf56d9ed84af3" +checksum = "cab112f0a86d568ea0e627cc1d6be74a1e9cd55214684db5561995f6dad897c6" +dependencies = [ + "num-traits", +] [[package]] name = "arbitrary" @@ -285,7 +363,7 @@ checksum = "0ae92a5119aa49cdbcf6b9f893fe4e1d98b04ccbf82ee0584ad948a44a734dea" dependencies = [ "proc-macro2", "quote", - "syn 2.0.118", + "syn 2.0.119", ] [[package]] @@ -294,15 +372,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" @@ -412,6 +481,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.119", +] + [[package]] name = "async-signal" version = "0.2.14" @@ -466,13 +546,13 @@ checksum = "8b75356056920673b02621b35afd0f7dda9306d03c79a30f5c56c44cf256e3de" [[package]] name = "async-trait" -version = "0.1.89" +version = "0.1.91" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "9035ad2d096bed7955a320ee7e2230574d28fd3c3a0f186cbea1ff3c7eed5dbb" +checksum = "ae36dc4177970ef04fde5178d3e2429882def40e57a451f919c098f72baa6cec" dependencies = [ "proc-macro2", "quote", - "syn 2.0.118", + "syn 3.0.0", ] [[package]] @@ -502,26 +582,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" @@ -545,6 +605,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" @@ -603,6 +687,12 @@ 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" @@ -611,18 +701,21 @@ checksum = "1e4b40c7323adcfc0a41c4b88143ed58346ff65a288fc144329c5c45e05d70c6" [[package]] name = "bitflags" -version = "2.13.0" +version = "1.3.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "bef38d45163c2f1dde094a7dfd33ccf595c92905c8f8f4fdc18d06fb1037718a" + +[[package]] +name = "bitflags" +version = "2.13.1" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "b4388bee8683e3d04af747c73422af53102d2bd24d9eadb6cbc100baef4b43f8" +checksum = "b588b76d00fde79687d7646a9b5bdf3cc0f655e0bbd080335a95d7e96f3587da" [[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" @@ -655,11 +748,21 @@ 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" +version = "0.7.7" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "5c0e531d93d39c34eef561e929e8a7f86d77a5af08aac4f6d6e39976c51858e9" +checksum = "56ed6191a7e78c36abdb16ab65341eefd73d64d303fffccdbb00d51e4205967b" [[package]] name = "bumpalo" @@ -696,9 +799,9 @@ dependencies = [ [[package]] name = "cc" -version = "1.2.67" +version = "1.3.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "e17dd265a7d0f31ef544e1b20e03add05d3b45b491b633b10d67145d2acc1a38" +checksum = "c89588d05638b5b4594a3348a2d6c20277e43a7f5c5202b05cc56888475a47b8" dependencies = [ "find-msvc-tools", "jobserver", @@ -706,6 +809,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" @@ -714,9 +827,9 @@ checksum = "9330f8b2ff13f34540b44e946ef35111825727b38d33286ef986142615121801" [[package]] name = "cfg_aliases" -version = "0.2.1" +version = "0.2.2" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "613afe47fcd5fac7ccf1db93babcb082c5994d996f20b8b159f2ad1658eb5724" +checksum = "f079e83a288787bcd14a6aea84cee5c87a67c5a3e660c30f557a3d24761b3527" [[package]] name = "chacha20" @@ -823,9 +936,9 @@ dependencies = [ [[package]] name = "clap" -version = "4.6.1" +version = "4.6.2" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "1ddb117e43bbf7dacf0a4190fef4d345b9bad68dfc649cb349e7d17d28428e51" +checksum = "dd059f9da4f5c36b3787f65d38ccaab1cc315f07b01f89abc8359ee6a8205011" dependencies = [ "clap_builder", "clap_derive", @@ -833,9 +946,9 @@ dependencies = [ [[package]] name = "clap_builder" -version = "4.6.0" +version = "4.6.2" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "714a53001bf66416adb0e2ef5ac857140e7dc3a0c48fb28b2f10762fc4b5069f" +checksum = "f09628afdcc538b57f3c6341e9c8e9970f18e4a481690a64974d7023bd33548b" dependencies = [ "anstream", "anstyle", @@ -852,7 +965,7 @@ dependencies = [ "heck 0.5.0", "proc-macro2", "quote", - "syn 2.0.118", + "syn 2.0.119", ] [[package]] @@ -861,6 +974,36 @@ 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.19", +] + +[[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 = "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" @@ -888,6 +1031,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" @@ -1002,7 +1155,7 @@ dependencies = [ "proc-macro2", "quote", "strsim", - "syn 2.0.118", + "syn 2.0.119", ] [[package]] @@ -1013,7 +1166,7 @@ checksum = "d38308df82d1080de0afee5d069fa14b0326a88c14f15c5ccda35b4a6c414c81" dependencies = [ "darling_core", "quote", - "syn 2.0.118", + "syn 2.0.119", ] [[package]] @@ -1022,6 +1175,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" @@ -1036,7 +1199,7 @@ checksum = "1e567bd82dcff979e4b03460c307b3cdc9e96fde3d73bed1496d2bc75d9dd62a" dependencies = [ "proc-macro2", "quote", - "syn 2.0.118", + "syn 2.0.119", ] [[package]] @@ -1090,7 +1253,7 @@ checksum = "1ac70aa55017e108007fbaf5aa0f54b021c98f92ff8af59d42eda9da96e3dd4f" dependencies = [ "proc-macro2", "quote", - "syn 2.0.118", + "syn 2.0.119", ] [[package]] @@ -1134,7 +1297,7 @@ checksum = "44f23cf4b44bfce11a86ace86f8a73ffdec849c9fd00a386a53d278bd9e81fb3" dependencies = [ "proc-macro2", "quote", - "syn 2.0.118", + "syn 2.0.119", ] [[package]] @@ -1193,7 +1356,7 @@ dependencies = [ "num-complex", "pulp", "rayon-core", - "smallvec", + "smallvec 1.15.2", "zune-inflate", ] @@ -1203,12 +1366,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" @@ -1218,6 +1375,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" @@ -1240,6 +1407,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" @@ -1259,11 +1441,17 @@ 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" +version = "0.3.33" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "8b147ee9d1f6d097cef9ce628cd2ee62288d963e16fb287bd9286455b241382d" +checksum = "a88cf1f829d945f548cf8fec32c61b1f202b6d93b45848602fc02af4b12ad218" dependencies = [ "futures-channel", "futures-core", @@ -1276,9 +1464,9 @@ dependencies = [ [[package]] name = "futures-channel" -version = "0.3.32" +version = "0.3.33" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "07bbe89c50d7a535e539b8c17bc0b49bdb77747034daa8087407d655f3f7cc1d" +checksum = "262590f4fe6afeb0bc83be1daa64e52657fe185690a958af7f3ad0e92085c5ae" dependencies = [ "futures-core", "futures-sink", @@ -1286,15 +1474,15 @@ dependencies = [ [[package]] name = "futures-core" -version = "0.3.32" +version = "0.3.33" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "7e3450815272ef58cec6d564423f6e755e25379b217b0bc688e295ba24df6b1d" +checksum = "2cd50c473c80f6d7c3670a752354b8e569b1a7cbfdc0419ec88e5edad85e0dc7" [[package]] name = "futures-executor" -version = "0.3.32" +version = "0.3.33" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "baf29c38818342a3b26b5b923639e7b1f4a61fc5e76102d4b1981c6dc7a7579d" +checksum = "6754879cc9f2c66f88c6e5c35344bb0bdb0708b0352b1201815667c7eabc7458" dependencies = [ "futures-core", "futures-task", @@ -1303,9 +1491,9 @@ dependencies = [ [[package]] name = "futures-io" -version = "0.3.32" +version = "0.3.33" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "cecba35d7ad927e23624b22ad55235f2239cfa44fd10428eecbeba6d6a717718" +checksum = "4577ecaa3c4f96589d473f679a71b596316f6641bc350038b962a5daf0085d7a" [[package]] name = "futures-lite" @@ -1322,26 +1510,26 @@ dependencies = [ [[package]] name = "futures-macro" -version = "0.3.32" +version = "0.3.33" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "e835b70203e41293343137df5c0664546da5745f82ec9b84d40be8336958447b" +checksum = "2d6d3cde68c518367be28956066ddfef33813991b77a55005a69dae04bf3b10b" dependencies = [ "proc-macro2", "quote", - "syn 2.0.118", + "syn 2.0.119", ] [[package]] name = "futures-sink" -version = "0.3.32" +version = "0.3.33" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "c39754e157331b013978ec91992bde1ac089843443c49cbc7f46150b0fad0893" +checksum = "e34418ac499d6305c2fb5ad0ed2f6ac998c5f8ca209b4510f7f94242c647e307" [[package]] name = "futures-task" -version = "0.3.32" +version = "0.3.33" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "037711b3d59c33004d3856fbdc83b99d4ff37a24768fa1be9ce3538a1cde4393" +checksum = "b231ed28831efb4a61a08580c4bc233ec56bc009f4cd8f52da2c3cb97df0c109" [[package]] name = "futures-timer" @@ -1351,9 +1539,9 @@ checksum = "af43fadb8a98512d547e37b4e92e0ced13e205c061b87b4623eff01d918d6968" [[package]] name = "futures-util" -version = "0.3.32" +version = "0.3.33" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "389ca41296e6190b48053de0321d02a77f32f8a5d2461dd38762c0593805c6d6" +checksum = "a77a90a256fce34da66415271e30f94ee91c57b04b8a2c042d9cf3220179deaa" dependencies = [ "futures-channel", "futures-core", @@ -1427,14 +1615,27 @@ 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", ] +[[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" @@ -1576,7 +1777,7 @@ dependencies = [ "httpdate", "itoa", "pin-project-lite", - "smallvec", + "smallvec 1.15.2", "tokio", "want", ] @@ -1594,7 +1795,7 @@ dependencies = [ "tokio", "tokio-rustls", "tower-service", - "webpki-roots 1.0.8", + "webpki-roots 1.0.9", ] [[package]] @@ -1681,7 +1882,7 @@ dependencies = [ "icu_normalizer_data", "icu_properties", "icu_provider", - "smallvec", + "smallvec 1.15.2", "zerovec", ] @@ -1739,7 +1940,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "3b0875f23caa03898994f6ddc501886a45c7d3d62d04d2d90788d47be1b1e4de" dependencies = [ "idna_adapter", - "smallvec", + "smallvec 1.15.2", "utf8_iter", ] @@ -1755,9 +1956,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", @@ -1765,7 +1966,6 @@ dependencies = [ "exr", "gif", "image-webp", - "moxcms", "num-traits", "png", "qoi", @@ -1787,6 +1987,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" @@ -1820,7 +2037,7 @@ checksum = "c34819042dc3d3971c46c2190835914dfbe0c3c13f61449b2997f4e9722dfa60" dependencies = [ "proc-macro2", "quote", - "syn 2.0.118", + "syn 2.0.119", ] [[package]] @@ -1837,9 +2054,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", ] @@ -1860,6 +2077,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" @@ -1965,6 +2188,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" @@ -2019,30 +2252,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" @@ -2058,6 +2319,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" @@ -2092,7 +2367,7 @@ checksum = "ed3955f1a9c7c0c15e092f9c887db08b1fc683305fdf6eb6684f22555355e202" dependencies = [ "proc-macro2", "quote", - "syn 2.0.118", + "syn 2.0.119", ] [[package]] @@ -2104,6 +2379,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" @@ -2122,6 +2407,18 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "071dfc062690e90b734c0b2273ce72ad0ffa95f0c74596bc250dcfd960262841" dependencies = [ "autocfg", + "libm", +] + +[[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]] @@ -2142,12 +2439,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.1", + "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.119", +] + +[[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" @@ -2161,10 +2535,23 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "57c0d7b74b563b49d38dae00a0c37d4d6de9b432382b2892f0574ddcae73fd0a" [[package]] -name = "pastey" -version = "0.1.1" +name = "pem" +version = "3.0.6" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "35fb2e5f958ec131621fdd531e9fc186ed768cbe395337403ae56c17a74c68ec" +checksum = "1d30c53c26bc5b31a98cd02d20f25a7c8567146caf63ed593a9d87b2775291be" +dependencies = [ + "base64", + "serde_core", +] + +[[package]] +name = "pem-rfc7468" +version = "1.0.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a6305423e0e7738146434843d1694d621cce767262b2a86910beab705e4493d9" +dependencies = [ + "base64ct", +] [[package]] name = "percent-encoding" @@ -2172,6 +2559,26 @@ 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.119", +] + [[package]] name = "pin-project-lite" version = "0.2.17" @@ -2195,13 +2602,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", @@ -2234,6 +2647,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" @@ -2260,9 +2688,9 @@ dependencies = [ [[package]] name = "proc-macro2" -version = "1.0.106" +version = "1.0.107" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "8fd00f0bb2e90d81d1044c2b32617f68fcb9fa3bb7640c23e9c748e53fb30934" +checksum = "985e7ec9bb745e6ce6535b544d84d6cd6f7ad8bd711c398938ae983b91a766d9" dependencies = [ "unicode-ident", ] @@ -2283,7 +2711,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "4488a4a36b9a4ba6b9334a32a39971f77c1436ec82c38707bce707699cc3bbcb" dependencies = [ "quote", - "syn 2.0.118", + "syn 2.0.119", ] [[package]] @@ -2309,12 +2737,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" @@ -2353,7 +2775,7 @@ dependencies = [ "rustc-hash", "rustls", "socket2", - "thiserror 2.0.18", + "thiserror 2.0.19", "tokio", "tracing", "web-time", @@ -2375,7 +2797,7 @@ dependencies = [ "rustls", "rustls-pki-types", "slab", - "thiserror 2.0.18", + "thiserror 2.0.19", "tinyvec", "tracing", "web-time", @@ -2397,9 +2819,9 @@ dependencies = [ [[package]] name = "quote" -version = "1.0.46" +version = "1.0.47" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "dfbc457d0c7a0759a614551b11a6409e5951f6c7537be1f1b7682b9ae9230368" +checksum = "1fbf4db142a473a8d80c26bbf18454ed458bf8d26c8219c331daecfdbd079001" dependencies = [ "proc-macro2", ] @@ -2492,6 +2914,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" @@ -2503,15 +2935,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", @@ -2526,21 +2956,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", @@ -2557,9 +2989,15 @@ version = "11.6.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "498cd0dc59d73224351ee52a95fee0f1a617a2eae0e7d9d720cc622c73a54186" dependencies = [ - "bitflags", + "bitflags 2.13.1", ] +[[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" @@ -2599,29 +3037,29 @@ dependencies = [ [[package]] name = "ref-cast" -version = "1.0.25" +version = "1.0.26" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "f354300ae66f76f1c85c5f84693f0ce81d747e2c3f21a45fef496d89c960bf7d" +checksum = "216e8f773d7923bcba9ceb86a86c93cabb3903a11872fc3f138c49630e50b96d" dependencies = [ "ref-cast-impl", ] [[package]] name = "ref-cast-impl" -version = "1.0.25" +version = "1.0.26" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "b7186006dcb21920990093f30e3dea63b7d6e977bf1256be20c3563a5db070da" +checksum = "2c9283685feec7d69af75fb0e858d5e7378f33fe4fc699383b2916ab9273e03c" dependencies = [ "proc-macro2", "quote", - "syn 2.0.118", + "syn 3.0.0", ] [[package]] name = "regex" -version = "1.13.0" +version = "1.13.1" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "2a0e75113e14dc5acb068cd0786884f214f1312650a3d36d269f5c4f3cdee8a2" +checksum = "f020237b6c8eed93db2e2cb53c00c60a8e1bc73da7d073199a1180401450218d" dependencies = [ "aho-corasick", "memchr", @@ -2631,9 +3069,9 @@ dependencies = [ [[package]] name = "regex-automata" -version = "0.4.15" +version = "0.4.16" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "1f388202e4b80542a0921078cc23b6333bcf1409c1e3f86404cae4766a6131db" +checksum = "8fcfdb36bda0c880c5931cdc7a2bcdc8ba4556847b9d912bca70bc94708711ad" dependencies = [ "aho-corasick", "memchr", @@ -2684,7 +3122,7 @@ dependencies = [ "wasm-bindgen-futures", "wasm-streams", "web-sys", - "webpki-roots 1.0.8", + "webpki-roots 1.0.9", ] [[package]] @@ -2703,7 +3141,7 @@ dependencies = [ "cfg-if", "getrandom 0.2.17", "libc", - "untrusted", + "untrusted 0.9.0", "windows-sys 0.52.0", ] @@ -2729,7 +3167,7 @@ dependencies = [ "serde", "serde_json", "sse-stream", - "thiserror 2.0.18", + "thiserror 2.0.19", "tokio", "tokio-stream", "tokio-util", @@ -2748,7 +3186,7 @@ dependencies = [ "proc-macro2", "quote", "serde_json", - "syn 2.0.118", + "syn 2.0.119", ] [[package]] @@ -2772,7 +3210,7 @@ dependencies = [ "proc-macro2", "quote", "rust-embed-utils", - "syn 2.0.118", + "syn 2.0.119", "walkdir", ] @@ -2798,7 +3236,7 @@ version = "0.38.44" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "fdb5bc1ae2baa591800df16c9ca78619bf65c0488b41b96ccec5d11220d8c154" dependencies = [ - "bitflags", + "bitflags 2.13.1", "errno", "libc", "linux-raw-sys 0.4.15", @@ -2811,7 +3249,7 @@ version = "1.1.4" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "b6fe4565b9518b83ef4f91bb47ce29620ca828bd32cb7e408f0062e9930ba190" dependencies = [ - "bitflags", + "bitflags 2.13.1", "errno", "libc", "linux-raw-sys 0.12.1", @@ -2824,6 +3262,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", @@ -2848,9 +3288,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]] @@ -2865,6 +3306,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" @@ -2874,6 +3324,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" @@ -2897,7 +3356,30 @@ dependencies = [ "proc-macro2", "quote", "serde_derive_internals", - "syn 2.0.118", + "syn 2.0.119", +] + +[[package]] +name = "security-framework" +version = "3.7.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b7f4bc775c73d9a02cde8bf7b2ec4c9d12743edf609006c7facc23998404cd1d" +dependencies = [ + "bitflags 2.13.1", + "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]] @@ -2908,9 +3390,9 @@ checksum = "8a7852d02fc848982e0c167ef163aaff9cd91dc640ba85e263cb1ce46fae51cd" [[package]] name = "serde" -version = "1.0.228" +version = "1.0.229" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "9a8e94ea7f378bd32cbbd37198a4a91436180c5bb472411e48b5ec2e2124ae9e" +checksum = "4148590afebada386688f18773da617792bf2ef03ffc1e4cbd2b1d45b023e0ba" dependencies = [ "serde_core", "serde_derive", @@ -2918,22 +3400,22 @@ dependencies = [ [[package]] name = "serde_core" -version = "1.0.228" +version = "1.0.229" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "41d385c7d4ca58e59fc732af25c3983b67ac852c1a25000afe1175de458b67ad" +checksum = "67dca2c9c51e58a4791a4b1ed58308b39c64224d349a935ab5039aa360942a48" dependencies = [ "serde_derive", ] [[package]] name = "serde_derive" -version = "1.0.228" +version = "1.0.229" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "d540f220d3187173da220f885ab66608367b6574e925011a9353e4badda91d79" +checksum = "e7a5d71263a5a7d47b41f6b3f06ba276f10cc18b0931f1799f710578e2309348" dependencies = [ "proc-macro2", "quote", - "syn 2.0.118", + "syn 3.0.0", ] [[package]] @@ -2944,7 +3426,7 @@ checksum = "18d26a20a969b9e3fdf2fc2d9f21eda6c40e2de84c9408bb5d3b05d499aae711" dependencies = [ "proc-macro2", "quote", - "syn 2.0.118", + "syn 2.0.119", ] [[package]] @@ -2971,6 +3453,24 @@ 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_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" @@ -2983,6 +3483,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" @@ -3032,11 +3545,24 @@ 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" +version = "0.3.10" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "703d5c7ef118737c72f1af64ad2f6f8c5e1921f818cdcb97b8fe6fc69bf66214" +checksum = "3a219298ac11a56ea9a6d2120044824d6f01aeb034955e7af7bc16858527deea" [[package]] name = "simd_helpers" @@ -3065,6 +3591,35 @@ 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 = "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.119", +] + [[package]] name = "socket2" version = "0.6.5" @@ -3075,6 +3630,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" @@ -3119,9 +3685,20 @@ dependencies = [ [[package]] name = "syn" -version = "2.0.118" +version = "2.0.119" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "872831b642d1a07999a962a351ed35b955ea2cfc8f3862091e2a240a84f17297" +dependencies = [ + "proc-macro2", + "quote", + "unicode-ident", +] + +[[package]] +name = "syn" +version = "3.0.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "1b9ae57f904213ebb649ce6895b8a66c66f0203b9319718f69a5612a065b1422" +checksum = "f2fac314a64dc9a36e61a9eb4261a5e9bbfbc922b27e518af97bc32b926cf967" dependencies = [ "proc-macro2", "quote", @@ -3145,9 +3722,39 @@ checksum = "728a70f3dbaf5bab7f0c4b1ac8d7ae5ea60a4b5549c8a5914361c99147a709d2" dependencies = [ "proc-macro2", "quote", - "syn 2.0.118", + "syn 2.0.119", ] +[[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" @@ -3172,11 +3779,11 @@ dependencies = [ [[package]] name = "thiserror" -version = "2.0.18" +version = "2.0.19" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "4288b5bcbc7920c07a1149a35cf9590a2aa808e0bc1eafaade0b80947865fbc4" +checksum = "09a43598840e33d5b0331f38c5e30d13bb11c11210a4b58f0d9b18a5a5eefcd9" dependencies = [ - "thiserror-impl 2.0.18", + "thiserror-impl 2.0.19", ] [[package]] @@ -3187,32 +3794,29 @@ checksum = "4fee6c4efc90059e10f81e6d42c60a18f76588c3d74cb83a0b242a2b6c7504c1" dependencies = [ "proc-macro2", "quote", - "syn 2.0.118", + "syn 2.0.119", ] [[package]] name = "thiserror-impl" -version = "2.0.18" +version = "2.0.19" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "ebc4ee7f67670e9b64d05fa4253e753e016c6c95ff35b89b7941d6b856dec1d5" +checksum = "43cbfe0cf76104d42a574802844187e84a305e531ed54455f11fbde0f10541cd" dependencies = [ "proc-macro2", "quote", - "syn 2.0.118", + "syn 3.0.0", ] [[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]] @@ -3272,9 +3876,9 @@ checksum = "1f3ccbac311fea05f86f61904b462b55fb3df8837a366dfc601a0161d0532f20" [[package]] name = "tokio" -version = "1.52.3" +version = "1.53.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "8fc7f01b389ac15039e4dc9531aa973a135d7a4135281b12d7c1bc79fd57fffe" +checksum = "d988bcd52dbe076d3d46903332f58c912b87a2c49b1428419a5845154762ffee" dependencies = [ "bytes", "libc", @@ -3288,13 +3892,13 @@ dependencies = [ [[package]] name = "tokio-macros" -version = "2.7.0" +version = "2.7.1" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "385a6cb71ab9ab790c5fe8d67f1645e6c450a7ce006a33de03daa956cf70a496" +checksum = "6328af13490e73a9b4694030fafd93f8c8c6a9dede33e821c3fc63eddf8042ba" dependencies = [ "proc-macro2", "quote", - "syn 2.0.118", + "syn 2.0.119", ] [[package]] @@ -3347,6 +3951,75 @@ 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 = "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" @@ -3369,7 +4042,7 @@ version = "0.6.11" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "4cfcf7e2740e6fc6d4d688b4ef00650406bb94adf4731e43c096c3a19fe40840" dependencies = [ - "bitflags", + "bitflags 2.13.1", "bytes", "futures-util", "http", @@ -3413,7 +4086,7 @@ checksum = "7490cfa5ec963746568740651ac6781f701c9c5ea257c58e057f3ba8cf69e8da" dependencies = [ "proc-macro2", "quote", - "syn 2.0.118", + "syn 2.0.119", ] [[package]] @@ -3431,6 +4104,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" @@ -3469,6 +4148,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" @@ -3487,6 +4172,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" @@ -3497,12 +4191,54 @@ 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.7.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a156c684c91ea7d62626509bce3cb4e1d9ed5c4d978f7b4352658f96a4c26b4a" + [[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" @@ -3528,6 +4264,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" @@ -3542,9 +4284,9 @@ checksum = "06abde3611657adf66d383f00b093d7faecc7fa57071cce2578660c9f1010821" [[package]] name = "uuid" -version = "1.23.5" +version = "1.24.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "ea5fab0d6c3c01ae70085a09cb03d4c7a1d6314e2b3e075392783396d724ca0a" +checksum = "bf3923a6f5c4c6382e0b653c4117f48d631ea17f38ed86e2a828e6f7412f5239" dependencies = [ "getrandom 0.4.3", "js-sys", @@ -3568,6 +4310,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" @@ -3650,7 +4404,7 @@ dependencies = [ "bumpalo", "proc-macro2", "quote", - "syn 2.0.118", + "syn 2.0.119", "wasm-bindgen-shared", ] @@ -3696,20 +4450,29 @@ 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" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "521bc38abb08001b01866da9f51eb7c5d647a19260e00054a8c7fd5f9e57f7a9" dependencies = [ - "webpki-roots 1.0.8", + "webpki-roots 1.0.9", ] [[package]] name = "webpki-roots" -version = "1.0.8" +version = "1.0.9" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "bf85cb06032201fa7c6f829d7db5a7e5aa45bcc0655327713065f6f0576731bf" +checksum = "7dcd9d09a39985f5344844e66b0c530a33843579125f23e21e9f0f220850f22a" dependencies = [ "rustls-pki-types", ] @@ -3744,6 +4507,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" @@ -3796,7 +4569,7 @@ checksum = "053e2e040ab57b9dc951b72c264860db7eb3b0200ba345b4e4c3b14f67855ddf" dependencies = [ "proc-macro2", "quote", - "syn 2.0.118", + "syn 2.0.119", ] [[package]] @@ -3807,7 +4580,7 @@ checksum = "3f316c4a2570ba26bbec722032c4099d8c8bc095efccdc15688708623367e358" dependencies = [ "proc-macro2", "quote", - "syn 2.0.118", + "syn 2.0.119", ] [[package]] @@ -3991,6 +4764,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" @@ -4020,10 +4802,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" @@ -4044,7 +4830,7 @@ checksum = "de844c262c8848816172cef550288e7dc6c7b7814b4ee56b3e1553f275f1858e" dependencies = [ "proc-macro2", "quote", - "syn 2.0.118", + "syn 2.0.119", "synstructure", ] @@ -4065,7 +4851,7 @@ checksum = "e2e817b7b52d0c7358d3246da9d69935ebb18116b2b102b4230dac079b4862f5" dependencies = [ "proc-macro2", "quote", - "syn 2.0.118", + "syn 2.0.119", ] [[package]] @@ -4085,7 +4871,7 @@ checksum = "11532158c46691caf0f2593ea8358fed6bbf68a0315e80aae9bd41fbade684a1" dependencies = [ "proc-macro2", "quote", - "syn 2.0.118", + "syn 2.0.119", "synstructure", ] @@ -4125,7 +4911,7 @@ checksum = "625dc425cab0dca6dc3c3319506e6593dcb08a9f387ea3b284dbd52a92c40555" dependencies = [ "proc-macro2", "quote", - "syn 2.0.118", + "syn 2.0.119", ] [[package]] @@ -4141,7 +4927,7 @@ dependencies = [ "flate2", "indexmap", "memchr", - "thiserror 2.0.18", + "thiserror 2.0.19", "zopfli", ] @@ -4165,9 +4951,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" @@ -4180,9 +4966,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 a136a13b..d4710d39 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -5,11 +5,13 @@ members = [ "crates/browser-driver", "crates/office", "crates/extension", + "crates/ocr", + "crates/science", ] resolver = "2" [workspace.package] -version = "0.1.1" +version = "0.1.2" edition = "2021" license = "MIT" repository = "https://github.com/A3S-Lab/Use" @@ -22,9 +24,14 @@ 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" +flate2 = "1" 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" @@ -32,11 +39,14 @@ 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"] } +tokio = { version = "1", features = ["fs", "io-std", "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"] } @@ -48,7 +58,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 +69,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 +77,7 @@ office = [ "dep:futures-util", "dep:getrandom", ] +ocr = ["dep:a3s-use-ocr"] extensions = ["dep:a3s-use-extension"] mcp = [ "dep:axum", @@ -81,10 +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-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 } @@ -108,5 +120,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 9b8376a1..20fa1e19 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 @@ -70,6 +71,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 @@ -95,7 +98,15 @@ 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 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 ``` Every domain argument accepted by `a3s use ...` can also be passed directly to @@ -103,8 +114,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, @@ -115,9 +126,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, @@ -129,13 +143,19 @@ 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**: Run pinned PP-OCRv6 detection and recognition + models locally through ONNX Runtime, with source digests and bounded layout + evidence +- **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 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; @@ -146,9 +166,11 @@ 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 | | 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 +179,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 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 | @@ -179,6 +202,8 @@ 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` | Local PP-OCRv6 engine, CLI, MCP tools, pinned models, and release-packaged Skill assets | +| `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 @@ -190,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 @@ -218,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" ``` @@ -242,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 @@ -279,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. @@ -361,9 +391,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 +436,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 +500,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 +593,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 @@ -515,12 +614,18 @@ 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 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. @@ -592,6 +697,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 +780,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 @@ -731,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 @@ -743,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 @@ -851,8 +966,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 +1128,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 +1256,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 +1541,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: @@ -1586,6 +1752,74 @@ 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`, +`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. 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, +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 --json +a3s use mcp serve ocr +``` + +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. + +## 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 @@ -1629,22 +1863,87 @@ 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`, `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 +1961,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 +2009,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 PP-OCRv6 ONNX CLI / MCP / Skill + + 0.1 compat + │ │ │ │ + └──────── capability snapshot/watch ───────────► A3S Code a3s-search ── Arc ──► a3s-use-browser @@ -1717,10 +2024,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 +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/browser-driver/Cargo.toml b/crates/browser-driver/Cargo.toml index c58f3d19..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,6 +13,7 @@ name = "a3s-use-browser-driver" path = "src/main.rs" [dependencies] +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-driver/skills/a3s-use-browser/SKILL.md b/crates/browser-driver/skills/a3s-use-browser/SKILL.md index 072d4ef5..1e0b537e 100644 --- a/crates/browser-driver/skills/a3s-use-browser/SKILL.md +++ b/crates/browser-driver/skills/a3s-use-browser/SKILL.md @@ -10,15 +10,24 @@ 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: +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/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/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/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/extension/Cargo.toml b/crates/extension/Cargo.toml index 93f63891..5a70b363 100644 --- a/crates/extension/Cargo.toml +++ b/crates/extension/Cargo.toml @@ -9,12 +9,22 @@ 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 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 ffc918a5..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,12 +20,19 @@ 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", "box", "capability", "office", + "office-compat", + "office-native", + "ocr", "capabilities", "component", "extension", @@ -437,7 +447,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/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/crates/ocr/Cargo.toml b/crates/ocr/Cargo.toml new file mode 100644 index 00000000..e9ba0dc7 --- /dev/null +++ b/crates/ocr/Cargo.toml @@ -0,0 +1,39 @@ +[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.2", path = "../core" } +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] +tempfile.workspace = true diff --git a/crates/ocr/README.md b/crates/ocr/README.md new file mode 100644 index 00000000..4a7f5251 --- /dev/null +++ b/crates/ocr/README.md @@ -0,0 +1,58 @@ +# 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 installing a separate extension. The native CLI and standard stdio MCP +share one local PP-OCRv6 implementation. + +There is one OCR provider: + +- provider: `pp-ocr-v6` +- engine: `onnx-runtime` +- model bundle: `PP-OCRv6_small` + +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 +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 --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 new file mode 100644 index 00000000..bbab0597 --- /dev/null +++ b/crates/ocr/skills/a3s-use-ocr/SKILL.md @@ -0,0 +1,54 @@ +--- +name: a3s-use-ocr +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`, `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. 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. +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. 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: + +```bash +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. + +## Boundaries + +- Only bounded local image files are accepted. URLs and PDF rasterization are + outside this domain. +- 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..e14ea852 --- /dev/null +++ b/crates/ocr/src/assets.rs @@ -0,0 +1,265 @@ +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, schemars::JsonSchema)] +#[serde(rename_all = "kebab-case")] +pub enum OcrInstallSource { + Environment, + Packaged, + Managed, + Missing, +} + +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize, schemars::JsonSchema)] +#[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( + "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())) +} + +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( + "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()) +} + +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 new file mode 100644 index 00000000..e5550740 --- /dev/null +++ b/crates/ocr/src/cli.rs @@ -0,0 +1,160 @@ +use std::path::PathBuf; + +use a3s_use_core::{UseError, UseResult}; +use clap::error::ErrorKind; +use clap::{Parser, Subcommand}; +use serde::Serialize; + +use crate::{OcrClient, OcrMcpServer, 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 local PP-OCRv6 readiness without reading an image. + Doctor, + /// Extract text and layout evidence from one local image. + Extract { path: PathBuf }, + /// Run an extension protocol surface. + Serve { + /// Serve standard MCP over stdin/stdout. + #[arg(long)] + mcp: bool, + }, +} + +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 } => { + CommandOutput::data(client.extract_with_first_use(OcrRequest { path }).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..7cf6f10a --- /dev/null +++ b/crates/ocr/src/client.rs @@ -0,0 +1,320 @@ +use std::path::{Path, PathBuf}; +use std::sync::{Arc, Mutex}; + +use a3s_use_core::{Artifact, Readiness, UseError, UseResult}; +use sha2::{Digest, Sha256}; +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, +}; +use crate::preprocess::decode_image; + +const MAX_INPUT_BYTES: u64 = 32 * 1024 * 1024; +const ENGINE_NAME: &str = "onnx-runtime"; + +#[derive(Clone)] +pub struct OcrClient { + loaded: Arc>>, +} + +struct LoadedEngine { + model_dir: PathBuf, + engine: PpOcrV6Engine, +} + +impl OcrClient { + pub fn from_env() -> UseResult { + Ok(Self { + loaded: Arc::new(Mutex::new(None)), + }) + } + + pub fn diagnostic(&self) -> OcrDiagnostic { + let status = ocr_status(); + let (readiness, suggestions) = if status.available { + (Readiness::Ready, Vec::new()) + } else if status.source == OcrInstallSource::Missing { + ( + Readiness::Missing, + vec![ + "Call the bounded ocr_install MCP tool, or run 'a3s install use/ocr' explicitly." + .to_string(), + ], + ) + } else { + ( + Readiness::Broken, + vec![ + "Call the bounded ocr_install MCP tool, or run 'a3s install use/ocr --force' explicitly." + .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 { + 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)?; + 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.", + ) + })?; + 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)?, + }); + } + let engine = loaded.as_mut().ok_or_else(|| { + UseError::new( + "use.ocr.runtime_failed", + "The local PP-OCRv6 engine failed to initialize.", + ) + })?; + let blocks = engine.engine.extract(&image)?; + build_result(source.artifact, blocks) + }) + .await + .map_err(|error| { + UseError::new( + "use.ocr.runtime_failed", + format!("The local PP-OCRv6 inference task failed: {error}"), + ) + })? + } +} + +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 { + 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 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 + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[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); + } + + #[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..17ecc8b8 --- /dev/null +++ b/crates/ocr/src/install.rs @@ -0,0 +1,789 @@ +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::{FirstUseInstallPolicy, 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, +} + +#[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, FirstUseInstallPolicy::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 { + 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 automatic_install_action( + status: &OcrRuntimeStatus, + policy: FirstUseInstallPolicy, +) -> 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 let Some(block) = policy.blocked_by() { + let reason = block.reason(); + 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 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) +} + +#[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), + FirstUseInstallPolicy::new(true, 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), + FirstUseInstallPolicy::new(false, false), + ) + .unwrap(); + + assert_eq!(action, AutoInstallAction::Install); + } + + #[test] + fn offline_and_no_auto_install_are_strict_boundaries() { + for policy in [ + FirstUseInstallPolicy::new(true, false), + FirstUseInstallPolicy::new(false, 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), + FirstUseInstallPolicy::new(false, false), + ) + .unwrap_err(); + + assert_eq!(error.code, "use.ocr.model_unreadable"); + } +} diff --git a/crates/ocr/src/lib.rs b/crates/ocr/src/lib.rs new file mode 100644 index 00000000..cc36b31f --- /dev/null +++ b/crates/ocr/src/lib.rs @@ -0,0 +1,29 @@ +//! 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. 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 postprocess; +mod preprocess; + +pub use assets::{ocr_status, OcrInstallSource, OcrRuntimeStatus}; +pub use client::OcrClient; +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, +}; + +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..229bb609 --- /dev/null +++ b/crates/ocr/src/mcp.rs @@ -0,0 +1,207 @@ +//! 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::{ + ensure_ppocr_v6_ready, OcrClient, OcrDiagnostic, OcrRequest, OcrResult, OcrRuntimeStatus, + 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 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, + 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_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", + 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_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. 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() + } + } +} + +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", "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 + .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(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/crates/ocr/src/models.rs b/crates/ocr/src/models.rs new file mode 100644 index 00000000..29b41198 --- /dev/null +++ b/crates/ocr/src/models.rs @@ -0,0 +1,100 @@ +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 { + PpOcrV6, +} + +#[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, +} + +#[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, + 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, + #[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 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 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")] + 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/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/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/office/skills/a3s-use-office/SKILL.md b/crates/office/skills/a3s-use-office/SKILL.md index 6fe1056f..37aa35bb 100644 --- a/crates/office/skills/a3s-use-office/SKILL.md +++ b/crates/office/skills/a3s-use-office/SKILL.md @@ -9,10 +9,28 @@ 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. +- 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, 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 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 +38,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,13 +72,19 @@ 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. 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. @@ -71,8 +97,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 +135,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 +170,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 +201,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 +223,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..dc7f791f 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 @@ -25,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 @@ -117,6 +123,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 +229,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 +324,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 +375,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 +514,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 @@ -504,6 +619,12 @@ 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 +In an A3S Code `use` worker, use an available +`mcp__use_office_compat__*` tool only when the native vocabulary lacks the +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/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("&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..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 @@ -37,6 +39,26 @@ 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 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. 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 +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 @@ -61,11 +83,19 @@ 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 -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 +103,18 @@ 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`; model readiness remains visible. +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 @@ -361,11 +399,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, @@ -404,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. @@ -417,7 +458,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 +520,16 @@ 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 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. +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: @@ -492,5 +543,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/docs/native-office.md b/docs/native-office.md index 6a0d2d0b..f5769357 100644 --- a/docs/native-office.md +++ b/docs/native-office.md @@ -68,10 +68,10 @@ sorting, filters, validation, conditional formatting, hyperlinks, drawings, images, charts, pivot tables and caches, slicers, sparklines, comments, OLE preservation, and CSV/TSV import. -The formula subsystem requires a real parser, dependency graph, recalculation -engine, dynamic-array spilling, reference rewriting, and a typed function -registry. Formula values are never evaluated by a shell or general-purpose -script runtime. +The formula subsystem uses a bounded parser, deterministic dependency graph, +native recalculation engine, dynamic-array spilling, reference rewriting, and +a typed closed function registry. Formula values are never evaluated by a +shell or general-purpose script runtime. ### Presentation @@ -199,9 +199,67 @@ paragraph/run/cell content and Presentation shape text; Spreadsheet cell paths upsert missing ordered rows and cells and maintain worksheet dimensions. Spreadsheet writes preserve explicit text, finite-number, boolean, and formula types. Formula writes strip an optional leading `=`, enforce Excel's 8192 -character bound, and mark the workbook for full recalculation; they do not yet -parse or evaluate formulas. Cell set/remove accepts normalized A1 rectangular -ranges of at most 100,000 cells and rolls back the whole operation on error. +character bound, parse the normalized body into a bounded source-spanned typed +AST, and mark the workbook for recalculation. Formula writes and import do not +implicitly calculate the workbook. The parser covers literals, Excel operator +precedence, calls and omitted arguments, parentheses and array constants, names +and structured references, qualified A1 cell/row/column references, and +range/intersection/union operators. Syntax errors report stable zero-based +UTF-8 byte and character offsets before mutation. + +The native calculation subsystem builds a deterministic graph across +worksheets, ranges, spills, and workbook- or worksheet-scoped names. +`NativeOfficeDocument::formula_dependency_graph`, +`calculate_spreadsheet_formulas`, and the registry-aware calculation variant +are read-only. The closed default 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`, together with ordinary operators, +typed errors, scalar/array broadcasting, spill references, and dynamic-array +results. + +`NativeOfficeEditor::recalculate_spreadsheet_formulas` and +`NativeOfficeMutation::RecalculateSpreadsheetFormulas` atomically calculate +and write typed cached values, canonical array anchors, spill children, and +calculated-workbook metadata. The mutation is exposed by replay, standard MCP +as `recalculate-spreadsheet-formulas`, and CLI as +`office native recalculate [--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/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/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..1f41f5bb 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; @@ -63,6 +64,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 +92,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 +109,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 +119,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 [--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 +138,7 @@ fn help() -> CommandOutput { "browser", "box", "office", + "ocr", "extension", "mcp" ] @@ -142,9 +150,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 +168,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 +236,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 +284,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 +296,7 @@ async fn component_list() -> UseResult { "browser".to_string(), "box".to_string(), "office".to_string(), + "ocr".to_string(), ]; human.extend( extensions @@ -398,6 +417,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")); @@ -433,15 +482,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) @@ -495,7 +607,28 @@ async fn component_uninstall(id: &str) -> UseResult { )); } } - if matches!(id, "browser" | "use/browser" | "office" | "use/office") { + 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" + ) { return Ok(CommandOutput::success( format!("No managed runtime files are owned for '{id}'."), serde_json::json!({ @@ -551,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; @@ -572,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)) } @@ -584,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, @@ -663,12 +811,15 @@ 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")] { + 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)) @@ -696,6 +847,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( @@ -858,6 +1024,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", } } @@ -883,11 +1051,22 @@ 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()), "box" | "use/box" => Some(crate::component_route::box_diagnostic()), "office" | "use/office" => Some(office_diagnostic()), + "ocr" | "use/ocr" => Some(ocr_diagnostic()), _ => None, } } @@ -919,9 +1098,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; } @@ -1031,7 +1217,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 +1244,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..23ff377b 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![ @@ -109,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/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/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 2c27ad51..c4d71b43 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -10,6 +10,10 @@ mod capability_registry; pub mod cli; mod component_route; mod extension_cli; +mod first_use; + +#[cfg(feature = "ocr")] +mod ocr_builtin; #[cfg(feature = "office")] mod office_artifact; @@ -37,5 +41,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..bfc6c201 100644 --- a/src/mcp/office.rs +++ b/src/mcp/office.rs @@ -84,9 +84,34 @@ 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" + 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 +133,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 +156,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 +182,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 +208,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 +241,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 +283,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 +364,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 +399,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 +436,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 +482,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 +515,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, @@ -464,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/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..7100d850 100644 --- a/src/mcp/office/tests.rs +++ b/src/mcp/office/tests.rs @@ -11,10 +11,10 @@ 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 = tools + let mut names: Vec<&str> = tools .iter() .map(|tool| tool.name.as_ref()) .collect::>(); @@ -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", @@ -36,6 +37,42 @@ 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)); + 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] @@ -339,6 +376,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/ocr_builtin.rs b/src/ocr_builtin.rs new file mode 100644 index 00000000..0e983c18 --- /dev/null +++ b/src/ocr_builtin.rs @@ -0,0 +1,93 @@ +//! 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.model_dir, + 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::PpOcrV6 => "pp-ocr-v6", + } +} + +#[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/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..9eae45b9 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"]) @@ -1126,6 +1146,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!({ @@ -1790,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() { @@ -1811,6 +1900,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() { @@ -1883,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") @@ -2100,6 +2218,109 @@ 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(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() { + 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() { 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/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() +} 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/"))); +}